From d909bf36ad0ea9992fb74f7280fc3a3df5e74601 Mon Sep 17 00:00:00 2001 From: hemanth1999k Date: Thu, 18 Jun 2026 19:02:30 -0500 Subject: [PATCH 1/2] Add hamming/blackman/bartlett window op converters MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit torch.hamming_window, torch.blackman_window and torch.bartlett_window had no converter — they failed with 'Unsupported fx node ... hamming_window' on the torch.export path (only hann_window was supported). Since window_length/periodic/alpha/beta are all known at conversion time, the (data-independent) window is materialized as a constant computed with the exact torch formulas. Handles the TorchScript and torch.export/ExecuTorch frontends (input-count differs), the periodic variants, and Hamming's alpha/beta overloads (hamming_window.periodic_alpha[_beta]) via torch_alias. Adds TestWindowFunctions. Verified against torch 2.9.0: convert + predict match PyTorch within fp16 tolerance across sizes, periodic settings, and custom alpha/beta. (kaiser_window is left out — it needs a Bessel I0 approximation.) --- .../converters/mil/frontend/torch/ops.py | 89 +++++++++++++++++++ .../mil/frontend/torch/test/test_torch_ops.py | 52 +++++++++++ 2 files changed, 141 insertions(+) diff --git a/coremltools/converters/mil/frontend/torch/ops.py b/coremltools/converters/mil/frontend/torch/ops.py index d2470423c..dffabfb98 100644 --- a/coremltools/converters/mil/frontend/torch/ops.py +++ b/coremltools/converters/mil/frontend/torch/ops.py @@ -8311,6 +8311,95 @@ def hann_window(context, node): context.add(sin_sq) +def _compute_window(kind: str, window_length: int, periodic: bool, + alpha: float = 0.54, beta: float = 0.46) -> np.ndarray: + """ + Materialize a torch window function (Hamming/Blackman/Bartlett) as a constant. + + window_length / periodic / alpha / beta are all known at conversion time, so the + (data-independent) window is computed here. torch builds a symmetric window of + length ``window_length + 1`` for the periodic case and drops the final sample. + """ + if window_length <= 0: + return np.zeros((0,), dtype=np.float32) + if window_length == 1: + return np.ones((1,), dtype=np.float32) + m = window_length + 1 if periodic else window_length + n = np.arange(m, dtype=np.float64) + denom = m - 1 + if kind == "hamming": + w = alpha - beta * np.cos(2.0 * np.pi * n / denom) + elif kind == "blackman": + w = 0.42 - 0.5 * np.cos(2.0 * np.pi * n / denom) + 0.08 * np.cos(4.0 * np.pi * n / denom) + elif kind == "bartlett": + w = 1.0 - np.abs(2.0 * n / denom - 1.0) + else: + raise ValueError(f"Unknown window kind: {kind}") + if periodic: + w = w[:-1] + return w.astype(np.float32) + + +def _parse_window_inputs(context, node, n_optional: int, defaults: list): + """ + Parse a torch window op into (window_length, [optional params]). + + The optional params are ``periodic`` (plus ``alpha``/``beta`` for Hamming). + TorchScript additionally carries the dtype/layout/device/pin_memory kwargs as the + trailing four graph inputs; torch.export / ExecuTorch pass only the positional + params, so we count the inputs accordingly. + """ + inputs = _get_inputs(context, node, min_expected=1) + if inputs[0].val is None: + raise NotImplementedError("variable 'window_length' not supported.") + window_length = int(inputs[0].val) + if context.frontend == TorchFrontend.TORCHSCRIPT: + n_present = len(inputs) - 1 - 4 # drop window_length + dtype/layout/device/pin_memory + else: + n_present = len(inputs) - 1 + # NOTE: builtin min/max are shadowed by the aten::min/max op handlers in this + # module, so clamp explicitly rather than calling min()/max(). + if n_present < 0: + n_present = 0 + elif n_present > n_optional: + n_present = n_optional + params = list(defaults) + for i in range(n_present): + var = inputs[1 + i] + if var is not None and var.val is not None: + params[i] = var.val + return window_length, params + + +@register_torch_op( + torch_alias=[ + "hamming_window.periodic", + "hamming_window.periodic_alpha", + "hamming_window.periodic_alpha_beta", + ] +) +def hamming_window(context, node): + window_length, (periodic, alpha, beta) = _parse_window_inputs( + context, node, n_optional=3, defaults=[True, 0.54, 0.46] + ) + window = _compute_window("hamming", window_length, bool(periodic), float(alpha), float(beta)) + context.add(mb.const(val=window, name=node.name)) + + +@register_torch_op(torch_alias=["blackman_window.periodic"]) +def blackman_window(context, node): + window_length, (periodic,) = _parse_window_inputs(context, node, n_optional=1, defaults=[True]) + window = _compute_window("blackman", window_length, bool(periodic)) + context.add(mb.const(val=window, name=node.name)) + + +@register_torch_op(torch_alias=["bartlett_window.periodic"]) +def bartlett_window(context, node): + window_length, (periodic,) = _parse_window_inputs(context, node, n_optional=1, defaults=[True]) + window = _compute_window("bartlett", window_length, bool(periodic)) + context.add(mb.const(val=window, name=node.name)) + + @register_torch_op def mse_loss(context, node): inputs = _get_inputs(context, node, expected=3) 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 47a6944c8..ed56b53a5 100644 --- a/coremltools/converters/mil/frontend/torch/test/test_torch_ops.py +++ b/coremltools/converters/mil/frontend/torch/test/test_torch_ops.py @@ -12593,6 +12593,58 @@ def forward(self, x): ) +class TestWindowFunctions(TorchBaseTest): + @pytest.mark.parametrize( + "compute_unit, backend, frontend, window_name, window_length, periodic", + itertools.product( + compute_units, + backends, + frontends, + ["hamming", "blackman", "bartlett"], + [1, 3, 6, 10, 12], + [True, False], + ), + ) + def test_window(self, compute_unit, backend, frontend, window_name, window_length, periodic): + window_op = { + "hamming": torch.hamming_window, + "blackman": torch.blackman_window, + "bartlett": torch.bartlett_window, + }[window_name] + + class WindowModel(nn.Module): + def forward(self, x): + return x + window_op(window_length, periodic=periodic) + + torch_in = torch.rand(window_length) + self.run_compare_torch( + torch_in, + WindowModel().eval(), + input_as_shape=False, + frontend=frontend, + backend=backend, + compute_unit=compute_unit, + ) + + @pytest.mark.parametrize( + "compute_unit, backend, frontend", + itertools.product(compute_units, backends, frontends), + ) + def test_hamming_window_alpha_beta(self, compute_unit, backend, frontend): + class HammingModel(nn.Module): + def forward(self, x): + return x + torch.hamming_window(16, periodic=True, alpha=0.5, beta=0.4) + + self.run_compare_torch( + torch.rand(16), + HammingModel().eval(), + input_as_shape=False, + frontend=frontend, + backend=backend, + compute_unit=compute_unit, + ) + + class TestTrace(TorchBaseTest): @pytest.mark.parametrize( "compute_unit, backend, frontend, shape", From 4b5891fa7978c25f62e42af76b32c6a2211430a7 Mon Sep 17 00:00:00 2001 From: hemanth1999k Date: Sat, 27 Jun 2026 18:09:04 -0500 Subject: [PATCH 2/2] Parametrize window test on op directly instead of name->op dict lookup --- .../mil/frontend/torch/test/test_torch_ops.py | 12 +++--------- 1 file changed, 3 insertions(+), 9 deletions(-) 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 ed56b53a5..b6ded4fc1 100644 --- a/coremltools/converters/mil/frontend/torch/test/test_torch_ops.py +++ b/coremltools/converters/mil/frontend/torch/test/test_torch_ops.py @@ -12595,23 +12595,17 @@ def forward(self, x): class TestWindowFunctions(TorchBaseTest): @pytest.mark.parametrize( - "compute_unit, backend, frontend, window_name, window_length, periodic", + "compute_unit, backend, frontend, window_op, window_length, periodic", itertools.product( compute_units, backends, frontends, - ["hamming", "blackman", "bartlett"], + [torch.hamming_window, torch.blackman_window, torch.bartlett_window], [1, 3, 6, 10, 12], [True, False], ), ) - def test_window(self, compute_unit, backend, frontend, window_name, window_length, periodic): - window_op = { - "hamming": torch.hamming_window, - "blackman": torch.blackman_window, - "bartlett": torch.bartlett_window, - }[window_name] - + def test_window(self, compute_unit, backend, frontend, window_op, window_length, periodic): class WindowModel(nn.Module): def forward(self, x): return x + window_op(window_length, periodic=periodic)