diff --git a/CHANGELOG.md b/CHANGELOG.md index 9ea3067..2aa7d1c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,7 +2,17 @@ ## [Unreleased] +### Fixed +- `hypyp.sync`: asking for the Metal backend on a metric that has no Metal kernel (PLV, CCorr, Coh, ImCoh, EnvCorr, PowCorr) used to run in NumPy silently, while the metric object reported `metal`. The computation still runs in NumPy, so no computed value changes, but it now warns and the metric object reports `numpy`. This holds for `optimization='metal'` and for a `priority` list that reaches Metal on a machine where Metal is available: the backends that follow in the list are not tried, which is left to 0.7.0. Only PLI, wPLI and ACCorr have a Metal kernel; the default selection of `optimization='auto'` is unchanged for the nine built-in metrics (#299) +- `hypyp.sync`: when a backend of a `priority` list can neither run on the machine nor be computed by the metric, and no other GPU backend of the list can be used, the warning now says so (for example "'plv' has no Metal implementation") instead of "No GPU backend available" +- `hypyp.sync`: a metric whose `_backend` was set to a backend it does not implement used to compute in NumPy without notice. It still computes in NumPy, now with a warning. A `_backend` that is not the name of a backend raises a `ValueError` naming the metric and the backends it implements; called through `compute_sync`, that error is still reworded as an unsupported metric (#306) +- Tests: the six Metal tests that compared NumPy with NumPy are replaced by tests of the fallback itself (warning, backend and result), the dispatch of every metric to every backend it implements is checked without a GPU, and the four tests that run a Metal kernel now assert that the Metal method was called (#300) + +### Added +- `BaseMetric.supports(backend)` tells whether a metric implements a backend, for example `PLI.supports('metal')` + ### Changed +- `hypyp.sync`: backend dispatch now lives in `BaseMetric.compute`, and each built-in metric implements `_compute_numpy` plus the optional `_compute_numba`, `_compute_torch`, `_compute_metal` and `_compute_cuda`. This is the recommended way to write a new metric: set `_dispatch_via_table = True` in the class. A subclass written the earlier way, which overrides `compute` and does its own dispatch, is still granted the backend it requests and is not subject to the capability check, whether it derives from `BaseMetric` or from a built-in metric. A subclass of a built-in metric that overrides `compute` and delegates to `super().compute()` keeps the earlier results: a backend without a `_compute_*` method computes in NumPy, now with a warning, and the CPU fallback of `optimization='auto'` still selects numba for it when numba is installed. A `_compute_*` method that a subclass adds for a backend its parent does not implement was never called before and is still not, unless the subclass sets `_dispatch_via_table = True` itself; a method of the parent that the subclass overrides is used, as before - The whole code base is formatted with `ruff format` (line length 88). The change is purely cosmetic: the syntax tree of every file is unchanged, apart from whitespace inside docstrings, and the Python examples of `hypyp/sync/README.md` are formatted too. The vendored `hypyp/ext` and the tutorial notebooks are left untouched. The formatting commit is listed in `.git-blame-ignore-revs`, so `git blame` skips it (run `git config blame.ignoreRevsFile .git-blame-ignore-revs` once in your clone; GitHub applies it automatically) - The CI now checks formatting (`ruff format --check`) and a small set of lint rules that only catch certain bugs: syntax errors, invalid comparisons and undefined names. Ruff comes from the new `lint` dependency group, which the `dev` group includes - `black` is removed from the `dev` dependency group, since `ruff format` replaces it diff --git a/hypyp/sync/__init__.py b/hypyp/sync/__init__.py index 1b0b419..10a1321 100644 --- a/hypyp/sync/__init__.py +++ b/hypyp/sync/__init__.py @@ -7,8 +7,8 @@ Public API ---------- ``BaseMetric`` - Abstract base class. Concrete metrics inherit from it and implement - ``BaseMetric.compute``. + Base class. Concrete metrics inherit from it and implement + ``_compute_numpy`` plus the optional accelerated ``_compute_*`` methods. Concrete metric classes (one per file): diff --git a/hypyp/sync/accorr.py b/hypyp/sync/accorr.py index dd83c2d..b4a9d64 100644 --- a/hypyp/sync/accorr.py +++ b/hypyp/sync/accorr.py @@ -79,6 +79,7 @@ class ACCorr(BaseMetric): """ name = "accorr" + _dispatch_via_table = True def __init__( self, @@ -109,16 +110,7 @@ def compute( con : np.ndarray ACCorr connectivity matrix with shape (n_epoch, n_freq, 2*n_ch, 2*n_ch). """ - if self._backend == "metal": - return self._compute_metal(complex_signal, n_samp, transpose_axes) - elif self._backend == "cuda_kernel": - return self._compute_cuda(complex_signal, n_samp, transpose_axes) - elif self._backend == "numba": - return self._compute_numba(complex_signal, n_samp, transpose_axes) - elif self._backend == "torch": - return self._compute_torch(complex_signal, n_samp, transpose_axes) - else: - return self._compute_numpy(complex_signal, n_samp, transpose_axes) + return super().compute(complex_signal, n_samp, transpose_axes) def _compute_metal( self, complex_signal: np.ndarray, n_samp: int, transpose_axes: tuple diff --git a/hypyp/sync/base.py b/hypyp/sync/base.py index 955d3f7..68cb1a3 100644 --- a/hypyp/sync/base.py +++ b/hypyp/sync/base.py @@ -8,9 +8,9 @@ This module is shared by every concrete metric in ``hypyp.sync``. It exposes: -- ``BaseMetric`` — abstract base. Concrete metrics override - ``BaseMetric.compute`` and rely on the shared backend-resolution and - warning-fallback logic. +- ``BaseMetric`` — base class. Concrete metrics implement the + ``_compute_*`` methods (``_compute_numpy`` at least) and rely on the + shared backend-resolution, warning-fallback and dispatch logic. - ``multiply_conjugate``, ``multiply_conjugate_time``, ``multiply_product`` — vectorised einsum kernels (numpy). - ``multiply_conjugate_torch``, ``multiply_conjugate_time_torch`` — @@ -43,7 +43,7 @@ """ import warnings -from abc import ABC, abstractmethod +from abc import ABC from typing import Optional import numpy as np @@ -279,17 +279,29 @@ def multiply_conjugate_time_torch(c, s): class BaseMetric(ABC): """ - Abstract base class for connectivity metrics. + Base class for connectivity metrics. - All connectivity metrics should inherit from this class and implement - the compute method. + A metric inherits from this class, sets ``_dispatch_via_table = True`` and + implements ``_compute_numpy`` plus any of the optional ``_compute_numba``, + ``_compute_torch``, ``_compute_metal`` and ``_compute_cuda``. Backend + selection, capability checks and dispatch are then handled here. + + A subclass of an existing metric that adds a ``_compute_*`` method for a + backend its parent does not implement must set ``_dispatch_via_table = + True`` itself for that method to be used. + + A subclass that overrides ``compute`` and does not set the flag itself + follows the earlier contract, whether it derives from this class or from a + built-in metric: it is granted whatever backend is requested and + available, and is itself responsible for honouring ``self._backend``. Parameters ---------- optimization : str, optional Optimization strategy for computation. Options: - None: standard numpy (default) - - 'auto': best available (torch > numba > numpy) + - 'auto': best backend for this metric and platform (see + ``_resolve_auto`` and ``AUTO_PRIORITY``) - 'numba': numba JIT compilation (falls back to numpy if unavailable) - 'torch': PyTorch with auto-detected GPU (falls back gracefully) @@ -303,6 +315,41 @@ class BaseMetric(ABC): name: str = "base" + #: Whether ``compute`` dispatches through ``_BACKEND_METHODS`` and backend + #: selection checks which ``_compute_*`` methods exist. ``None`` means + #: "inferred": yes, unless the class overrides ``compute``. The built-in + #: metrics override ``compute`` only to carry a docstring, so they set + #: the flag explicitly; so should a new metric written the same way. + #: The flag vouches for the ``compute`` of the class that sets it: a + #: descendant that overrides ``compute`` again may add backends of its + #: own there, so it is no longer capability-checked unless it sets the + #: flag too. Setting it back to ``None`` in a descendant does not restore + #: the inference: the nearest class of the MRO that sets ``True`` or + #: ``False`` decides. + _dispatch_via_table: Optional[bool] = None + + #: Maps a backend name to the method implementing it. This table is the + #: single source of truth for dispatch: ``compute`` looks the backend up + #: here, so an unrecognised backend raises ``ValueError`` instead of silently + #: falling through to numpy, and ``supports`` derives capability from the + #: methods a subclass actually defines rather than from a hand-kept list. + _BACKEND_METHODS = { + "numpy": "_compute_numpy", + "numba": "_compute_numba", + "torch": "_compute_torch", + "metal": "_compute_metal", + "cuda_kernel": "_compute_cuda", + } + + #: Human-readable backend names, used in fallback warnings. + _BACKEND_LABELS = { + "numpy": "numpy", + "numba": "numba", + "torch": "torch", + "metal": "Metal", + "cuda_kernel": "CUDA", + } + def __init__( self, optimization: Optional[str] = None, priority: Optional[list] = None ): @@ -310,6 +357,156 @@ def __init__( self._priority = priority self._backend, self._device = self._resolve_optimization(optimization, priority) + @classmethod + def supports(cls, backend: str) -> bool: + """ + Whether this metric implements ``backend``. + + Capability is derived from the presence of the corresponding + ``_compute_*`` method, so it cannot drift out of sync with the code. + Not every metric has every backend — Metal kernels exist only for the + sign-based metrics and ACCorr, because torch on MPS is faster for the + einsum metrics at every channel count (see ``AUTO_PRIORITY``). + + A subclass written against the earlier contract, which overrides + ``compute`` and branches on ``self._backend`` itself, cannot be + inspected this way. Unless it sets ``_dispatch_via_table`` itself, it + is trusted with every known backend, exactly as before the capability + check existed. This holds for a subclass of a built-in metric too. + + Parameters + ---------- + backend : str + One of ``'numpy'``, ``'numba'``, ``'torch'``, ``'metal'``, + ``'cuda_kernel'``. An unknown name returns ``False``. + + Returns + ------- + bool + True if the metric can run on ``backend``. + + Examples + -------- + >>> from hypyp.sync import PLI, PLV + >>> PLI.supports('metal'), PLV.supports('metal') + (True, False) + """ + method = cls._BACKEND_METHODS.get(backend) + if method is None: + return False + if not cls._checks_capability(): + return True + return cls._implements(backend) + + @classmethod + def _implements(cls, backend: str) -> bool: + """Whether the class has a real ``_compute_*`` method for ``backend``. + + The default ``_compute_numpy`` of this class only raises, and a + non-callable attribute is a placeholder: neither is an implementation. + + A backend also counts only if the class that adopted the table + dispatch (the one that sets ``_dispatch_via_table``) already had a + method for it. Before the table, the ``compute`` of each built-in + metric called a fixed set of ``_compute_*`` methods: a descendant + could override one of them, but a method it added for another backend + was never called. The 0.6 series changes no computed value, so such a + method stays unused until the descendant sets the flag itself. + """ + + def is_implementation(candidate) -> bool: + return callable(candidate) and ( + candidate is not BaseMetric.__dict__["_compute_numpy"] + ) + + method = cls._BACKEND_METHODS.get(backend) + if not method or not is_implementation(getattr(cls, method, None)): + return False + owner, _ = cls._dispatch_owner() + if owner is None: + return True + # The flag may sit on a mixin: the class that adopted the dispatch is + # then the metric class that brought the mixin in. + adopters = [ + k for k in cls.__mro__ if issubclass(k, BaseMetric) and owner in k.__mro__ + ] + mro = cls.__mro__ + for klass in mro[mro.index(adopters[-1]) :]: + if method in klass.__dict__: + # getattr, not __dict__: a classmethod or staticmethod is + # callable only once its descriptor is resolved. + return is_implementation(getattr(klass, method)) + return False + + @classmethod + def _dispatch_owner(cls) -> tuple: + """The nearest class of the MRO that sets ``_dispatch_via_table``, and + the value it sets; ``(None, None)`` when no class does.""" + for klass in cls.__mro__: + flag = klass.__dict__.get("_dispatch_via_table") + if flag is not None: + return klass, flag + return None, None + + @classmethod + def _dispatches_via_table(cls) -> bool: + """Whether ``BaseMetric.compute`` dispatches for this class. + + Explicit when a class of the MRO sets ``_dispatch_via_table``. + Otherwise inferred: a class that overrides ``compute`` is taken to do + its own dispatch, as the contract was before ``compute`` became + concrete. + """ + owner, flag = cls._dispatch_owner() + if owner is None: + return cls.compute is BaseMetric.compute + return flag + + @classmethod + def _checks_capability(cls) -> bool: + """Whether backend selection may rely on the ``_compute_*`` methods. + + True when the table dispatch is the only dispatch: the class uses it, + and ``compute`` has not been overridden below the class that set the + flag. A descendant that overrides ``compute`` again may handle + backends there that no ``_compute_*`` method reveals. + """ + if not cls._dispatches_via_table(): + return False + owner, _ = cls._dispatch_owner() + if owner is None: + return True + # The flag vouches for the compute that its owner resolves to, which + # the owner need not define itself (a mixin, or a metric that keeps + # the compute of this class). Compare the functions rather than the + # classes that hold them: a descendant that rebinds the very same + # function (``compute = PLV.compute``) has not changed the dispatch. + mro = cls.__mro__ + + def first_compute(classes: tuple): + # A flag owner placed after every class that defines compute (a + # mixin listed after BaseMetric) vouches for the table dispatch. + return next( + (k.__dict__["compute"] for k in classes if "compute" in k.__dict__), + BaseMetric.__dict__["compute"], + ) + + return first_compute(mro) is first_compute(mro[mro.index(owner) :]) + + @classmethod + def _cpu_fallback(cls) -> tuple: + """CPU backend used when no GPU backend can be selected: numba when it + is installed and the metric supports it, numpy otherwise. + + A class that overrides ``compute`` is trusted with numba here, as it + was before the capability check: it may handle numba in its own + ``compute``. If it only delegates and no ``_compute_numba`` exists, + the dispatch computes in numpy with a warning. + """ + if NUMBA_AVAILABLE and cls.supports("numba"): + return "numba", "cpu" + return "numpy", "cpu" + @classmethod def _resolve_optimization( cls, optimization: Optional[str] = None, priority: Optional[list] = None @@ -355,8 +552,11 @@ def _resolve_optimization( ----- Fallback cascade for ``'auto'`` (per-metric, per-platform): Iterates ``AUTO_PRIORITY[metric][platform]`` and returns the - first available backend. Falls back to numba → numpy if no - GPU backend is available. + first available backend the metric implements. An available + backend the metric does not implement ends the search in numpy + with a warning, unless the metric has no numpy implementation, + in which case the search goes on. Falls back to numba → numpy if no GPU backend + is available. Fallback cascade for explicit backends when unavailable: requested backend → numpy (with UserWarning) @@ -367,6 +567,27 @@ def _resolve_optimization( if optimization == "auto": return cls._resolve_auto(priority) + if optimization not in ("numba", "torch", "metal", "cuda_kernel"): + raise ValueError( + f"Unknown optimization '{optimization}'. " + f"Options: None, 'auto', 'numba', 'torch', 'metal', 'cuda_kernel'" + ) + + # Capability before availability: a backend the machine can run is + # still useless if this metric has no implementation for it. Without + # this check the backend was accepted and dispatch quietly returned a + # numpy result — the caller believed they were on the GPU. + if not cls.supports(optimization): + label = cls._BACKEND_LABELS[optimization] + warnings.warn( + f"{cls.name!r} has no {label} implementation, falling back to " + f"numpy. Use optimization='auto' to select the best backend " + f"available for this metric.", + UserWarning, + stacklevel=3, + ) + return "numpy", "cpu" + if optimization == "numba": if NUMBA_AVAILABLE: return "numba", "cpu" @@ -411,6 +632,9 @@ def _resolve_optimization( ) return "numpy", "cpu" + # Unreachable: the membership test above already rejected any other + # value. Kept as a guard in case a backend is added to that tuple + # without a matching branch here. raise ValueError( f"Unknown optimization '{optimization}'. " f"Options: None, 'auto', 'numba', 'torch', 'metal', 'cuda_kernel'" @@ -423,7 +647,10 @@ def _resolve_auto(cls, priority: Optional[list] = None) -> tuple: Uses the ``AUTO_PRIORITY`` table compiled from Mac M4 Max and Narval A100 benchmarks. Iterates the priority list and returns - the first available backend. + the first available backend the metric implements. An available + backend the metric does not implement ends the search in numpy with + a warning, unless the metric has no numpy implementation (see the + comment in the loop). Parameters ---------- @@ -455,14 +682,43 @@ def _resolve_auto(cls, priority: Optional[list] = None) -> tuple: UserWarning, stacklevel=4, ) - if NUMBA_AVAILABLE: - return "numba", "cpu" - return "numpy", "cpu" + return cls._cpu_fallback() if priority is None: priority = AUTO_PRIORITY.get(cls.name, {}).get(platform, []) + # Backends of the priority list this metric has no implementation for, + # remembered so the fallback warning can give the real reason. + unimplemented = [] + available = { + "torch": TORCH_AVAILABLE, + "metal": METAL_AVAILABLE, + "cuda_kernel": CUPY_AVAILABLE, + } for backend in priority: + if not cls.supports(backend): + label = cls._BACKEND_LABELS.get(backend) + # Earlier versions selected an available backend here even + # though the metric has no implementation for it, and the + # computation then ran in numpy without notice. The 0.6 + # series does not change computed values, so the selection + # still ends in numpy, now with a warning. Moving on to the + # next backend of the list instead is left to 0.7.0. + # (A metric without a numpy implementation could not exist + # in those versions, so for it the search simply goes on.) + if label and available.get(backend) and cls._implements("numpy"): + warnings.warn( + f"{cls.name!r} has no {label} implementation: computing " + f"with numpy, as earlier versions did silently. The " + f"backends that follow in the priority list are not " + f"tried.", + UserWarning, + stacklevel=4, + ) + return "numpy", "cpu" + if label: + unimplemented.append(label) + continue if backend == "torch" and TORCH_AVAILABLE: return cls._resolve_torch() if backend == "metal" and METAL_AVAILABLE: @@ -470,16 +726,25 @@ def _resolve_auto(cls, priority: Optional[list] = None) -> tuple: if backend == "cuda_kernel" and CUPY_AVAILABLE: return "cuda_kernel", "cuda" - # No GPU backend from priority list available — fall back + # No GPU backend from priority list available — fall back. When a + # backend was skipped for lack of an implementation, say so: "no GPU + # backend available" would be false on a machine that has one. + if unimplemented: + reason = ( + f"{cls.name!r} has no {' or '.join(unimplemented)} " + f"implementation, and no GPU backend of the priority list " + f"is available on platform '{platform}'." + ) + else: + reason = ( + f"No GPU backend available for {cls.name!r} on platform '{platform}'." + ) warnings.warn( - f"No GPU backend available for {cls.name!r} on platform " - f"'{platform}'. Falling back to CPU.", + f"{reason} Falling back to CPU.", UserWarning, stacklevel=4, ) - if NUMBA_AVAILABLE: - return "numba", "cpu" - return "numpy", "cpu" + return cls._cpu_fallback() @staticmethod def _resolve_torch() -> tuple: @@ -513,12 +778,115 @@ def _resolve_torch() -> tuple: warnings.warn("No GPU found, using torch on CPU", UserWarning, stacklevel=4) return "torch", "cpu" - @abstractmethod def compute( self, complex_signal: np.ndarray, n_samp: int, transpose_axes: tuple ) -> np.ndarray: """ - Compute the connectivity metric. + Compute the connectivity metric on the resolved backend. + + Dispatch is table-driven via ``_BACKEND_METHODS``: the backend chosen at + construction selects the ``_compute_*`` method to run. Subclasses + normally implement those methods and leave this one alone; a subclass + that overrides ``compute`` itself bypasses the dispatch and is + responsible for honouring ``self._backend``. + + Parameters + ---------- + complex_signal : np.ndarray + Complex analytic signals with shape (n_epochs, n_freq, 2*n_channels, n_times). + n_samp : int + Number of time samples. + transpose_axes : tuple + Axes to transpose for matrix multiplication. + + Returns + ------- + con : np.ndarray + Connectivity matrix with shape (n_epoch, n_freq, 2*n_ch, 2*n_ch). + + Raises + ------ + NotImplementedError + If the backend is numpy and the metric has no ``_compute_numpy`` + (a metric may implement an accelerated backend alone, but then + cannot serve the default ``optimization=None``). + ValueError + If ``self._backend`` is not the name of a backend, or is a backend + the metric has no ``_compute_*`` method for while it has no numpy + implementation either. The message names the metric and the + backends it does implement. + + Warns + ----- + UserWarning + If ``self._backend`` is a known backend the metric has no + ``_compute_*`` method for. The computation then runs in numpy, as + the earlier hand-written ``if/elif`` chain of each metric did + without notice. Backend selection never produces this state for a + built-in metric; it arises when ``_backend`` is set by hand, or in + a subclass that overrides ``compute`` and delegates here. + + Notes + ----- + Output precision depends on the backend and on the metric. The Metal + kernels and torch on MPS compute in ``float32`` whatever the input; + for the other combinations see each metric. + + A subclass of ``BaseMetric`` that does its own dispatch (see + ``_dispatch_via_table``) may call ``super().compute(...)`` from its own + ``compute``. The base method was then abstract with an empty body and + returned ``None``; it still does for such a subclass. A subclass of a + built-in metric that overrides ``compute`` and delegates to + ``super().compute(...)`` gets the table dispatch, which looks the + method up on the instance, so a ``_compute_*`` method it overrides is + used. A method it adds for a backend its parent does not implement is + used only if the subclass sets ``_dispatch_via_table`` itself; + otherwise the warning or errors above apply to that backend. + """ + if not self._dispatches_via_table(): + # Reached through super().compute() from a subclass that does its + # own dispatch: behave as the former abstract method did. + return None + if not self._implements(self._backend): + if self._backend == "numpy": + raise NotImplementedError( + f"{type(self).__name__} must implement _compute_numpy " + f"(or override compute)." + ) + implemented = [b for b in self._BACKEND_METHODS if self._implements(b)] + if self._backend in self._BACKEND_METHODS and "numpy" in implemented: + # The per-metric if/elif chains this dispatch replaces ended + # in the numpy implementation. The 0.6 series changes no + # computed value, so a known backend without a method still + # computes in numpy, now with a warning. + warnings.warn( + f"{self.name!r} has no " + f"{self._BACKEND_LABELS.get(self._backend, self._backend)} " + f"implementation: computing with numpy, as earlier " + f"versions did silently.", + UserWarning, + stacklevel=2, + ) + return self._compute_numpy(complex_signal, n_samp, transpose_axes) + raise ValueError( + f"{self.name!r} cannot run on backend {self._backend!r}. " + f"Backends implemented for this metric: {implemented}." + ) + method = getattr(self, self._BACKEND_METHODS[self._backend]) + return method(complex_signal, n_samp, transpose_axes) + + def _compute_numpy( + self, complex_signal: np.ndarray, n_samp: int, transpose_axes: tuple + ) -> np.ndarray: + """ + Reference implementation, in numpy. + + Every metric should provide this: it is the correctness oracle the + accelerated backends are validated against, and the fallback target + whenever a requested backend is unavailable or unimplemented. It is + deliberately not an abstract method, so that a subclass written + against the earlier contract (its own ``compute``) can still be + instantiated; this default raises ``NotImplementedError``. Parameters ---------- @@ -534,4 +902,7 @@ def compute( con : np.ndarray Connectivity matrix with shape (n_epoch, n_freq, 2*n_ch, 2*n_ch). """ - pass + raise NotImplementedError( + f"{type(self).__name__} must implement _compute_numpy " + f"(or override compute)." + ) diff --git a/hypyp/sync/ccorr.py b/hypyp/sync/ccorr.py index b859776..262745e 100644 --- a/hypyp/sync/ccorr.py +++ b/hypyp/sync/ccorr.py @@ -38,6 +38,7 @@ class CCorr(BaseMetric): """ name = "ccorr" + _dispatch_via_table = True def compute( self, complex_signal: np.ndarray, n_samp: int, transpose_axes: tuple @@ -61,13 +62,7 @@ def compute( con : np.ndarray CCorr connectivity matrix with shape (n_epoch, n_freq, 2*n_ch, 2*n_ch). """ - if self._backend == "cuda_kernel": - return self._compute_cuda(complex_signal, n_samp, transpose_axes) - elif self._backend == "torch": - return self._compute_torch(complex_signal, n_samp, transpose_axes) - elif self._backend == "numba": - return self._compute_numba(complex_signal, n_samp, transpose_axes) - return self._compute_numpy(complex_signal, n_samp, transpose_axes) + return super().compute(complex_signal, n_samp, transpose_axes) def _compute_cuda(self, complex_signal, n_samp, transpose_axes): """CUDA kernel for CCorr.""" diff --git a/hypyp/sync/coh.py b/hypyp/sync/coh.py index 3cd37bf..ee2b334 100644 --- a/hypyp/sync/coh.py +++ b/hypyp/sync/coh.py @@ -44,6 +44,7 @@ class Coh(BaseMetric): """ name = "coh" + _dispatch_via_table = True def compute( self, complex_signal: np.ndarray, n_samp: int, transpose_axes: tuple @@ -67,13 +68,7 @@ def compute( con : np.ndarray Coherence connectivity matrix with shape (n_epoch, n_freq, 2*n_ch, 2*n_ch). """ - if self._backend == "cuda_kernel": - return self._compute_cuda(complex_signal, n_samp, transpose_axes) - elif self._backend == "torch": - return self._compute_torch(complex_signal, n_samp, transpose_axes) - elif self._backend == "numba": - return self._compute_numba(complex_signal, n_samp, transpose_axes) - return self._compute_numpy(complex_signal, n_samp, transpose_axes) + return super().compute(complex_signal, n_samp, transpose_axes) def _compute_cuda(self, complex_signal, n_samp, transpose_axes): """CUDA kernel for Coherence.""" diff --git a/hypyp/sync/envelope_corr.py b/hypyp/sync/envelope_corr.py index 234c611..45da243 100644 --- a/hypyp/sync/envelope_corr.py +++ b/hypyp/sync/envelope_corr.py @@ -47,6 +47,7 @@ class EnvCorr(BaseMetric): """ name = "envcorr" + _dispatch_via_table = True def compute( self, complex_signal: np.ndarray, n_samp: int, transpose_axes: tuple @@ -71,13 +72,7 @@ def compute( Envelope Correlation connectivity matrix with shape (n_epoch, n_freq, 2*n_ch, 2*n_ch). """ - if self._backend == "cuda_kernel": - return self._compute_cuda(complex_signal, n_samp, transpose_axes) - elif self._backend == "torch": - return self._compute_torch(complex_signal, n_samp, transpose_axes) - elif self._backend == "numba": - return self._compute_numba(complex_signal, n_samp, transpose_axes) - return self._compute_numpy(complex_signal, n_samp, transpose_axes) + return super().compute(complex_signal, n_samp, transpose_axes) def _compute_cuda(self, complex_signal, n_samp, transpose_axes): """CUDA kernel for Envelope Correlation.""" diff --git a/hypyp/sync/imaginary_coh.py b/hypyp/sync/imaginary_coh.py index 630d4c5..54a5b07 100644 --- a/hypyp/sync/imaginary_coh.py +++ b/hypyp/sync/imaginary_coh.py @@ -46,6 +46,7 @@ class ImCoh(BaseMetric): """ name = "imcoh" + _dispatch_via_table = True def compute( self, complex_signal: np.ndarray, n_samp: int, transpose_axes: tuple @@ -70,13 +71,7 @@ def compute( Imaginary Coherence connectivity matrix with shape (n_epoch, n_freq, 2*n_ch, 2*n_ch). """ - if self._backend == "cuda_kernel": - return self._compute_cuda(complex_signal, n_samp, transpose_axes) - elif self._backend == "torch": - return self._compute_torch(complex_signal, n_samp, transpose_axes) - elif self._backend == "numba": - return self._compute_numba(complex_signal, n_samp, transpose_axes) - return self._compute_numpy(complex_signal, n_samp, transpose_axes) + return super().compute(complex_signal, n_samp, transpose_axes) def _compute_cuda(self, complex_signal, n_samp, transpose_axes): """CUDA kernel for Imaginary Coherence.""" diff --git a/hypyp/sync/pli.py b/hypyp/sync/pli.py index 669b88c..eec2f96 100644 --- a/hypyp/sync/pli.py +++ b/hypyp/sync/pli.py @@ -48,6 +48,7 @@ class PLI(BaseMetric): """ name = "pli" + _dispatch_via_table = True def compute( self, complex_signal: np.ndarray, n_samp: int, transpose_axes: tuple @@ -71,15 +72,7 @@ def compute( con : np.ndarray PLI connectivity matrix with shape (n_epoch, n_freq, 2*n_ch, 2*n_ch). """ - if self._backend == "metal": - return self._compute_metal(complex_signal, n_samp, transpose_axes) - elif self._backend == "cuda_kernel": - return self._compute_cuda(complex_signal, n_samp, transpose_axes) - elif self._backend == "torch": - return self._compute_torch(complex_signal, n_samp, transpose_axes) - elif self._backend == "numba": - return self._compute_numba(complex_signal, n_samp, transpose_axes) - return self._compute_numpy(complex_signal, n_samp, transpose_axes) + return super().compute(complex_signal, n_samp, transpose_axes) def _compute_numpy( self, complex_signal: np.ndarray, n_samp: int, transpose_axes: tuple diff --git a/hypyp/sync/plv.py b/hypyp/sync/plv.py index 7fc6194..682b3b2 100644 --- a/hypyp/sync/plv.py +++ b/hypyp/sync/plv.py @@ -44,6 +44,7 @@ class PLV(BaseMetric): """ name = "plv" + _dispatch_via_table = True def compute( self, complex_signal: np.ndarray, n_samp: int, transpose_axes: tuple @@ -67,13 +68,7 @@ def compute( con : np.ndarray PLV connectivity matrix with shape (n_epoch, n_freq, 2*n_ch, 2*n_ch). """ - if self._backend == "cuda_kernel": - return self._compute_cuda(complex_signal, n_samp, transpose_axes) - elif self._backend == "torch": - return self._compute_torch(complex_signal, n_samp, transpose_axes) - elif self._backend == "numba": - return self._compute_numba(complex_signal, n_samp, transpose_axes) - return self._compute_numpy(complex_signal, n_samp, transpose_axes) + return super().compute(complex_signal, n_samp, transpose_axes) def _compute_cuda(self, complex_signal, n_samp, transpose_axes): """CUDA kernel for PLV.""" diff --git a/hypyp/sync/pow_corr.py b/hypyp/sync/pow_corr.py index 21d403c..bfffcc1 100644 --- a/hypyp/sync/pow_corr.py +++ b/hypyp/sync/pow_corr.py @@ -49,6 +49,7 @@ class PowCorr(BaseMetric): """ name = "powcorr" + _dispatch_via_table = True def compute( self, complex_signal: np.ndarray, n_samp: int, transpose_axes: tuple @@ -73,13 +74,7 @@ def compute( Power Correlation connectivity matrix with shape (n_epoch, n_freq, 2*n_ch, 2*n_ch). """ - if self._backend == "cuda_kernel": - return self._compute_cuda(complex_signal, n_samp, transpose_axes) - elif self._backend == "torch": - return self._compute_torch(complex_signal, n_samp, transpose_axes) - elif self._backend == "numba": - return self._compute_numba(complex_signal, n_samp, transpose_axes) - return self._compute_numpy(complex_signal, n_samp, transpose_axes) + return super().compute(complex_signal, n_samp, transpose_axes) def _compute_cuda(self, complex_signal, n_samp, transpose_axes): """CUDA kernel for Power Correlation.""" diff --git a/hypyp/sync/wpli.py b/hypyp/sync/wpli.py index a42ca2d..7f0f627 100644 --- a/hypyp/sync/wpli.py +++ b/hypyp/sync/wpli.py @@ -47,6 +47,7 @@ class WPLI(BaseMetric): """ name = "wpli" + _dispatch_via_table = True def compute( self, complex_signal: np.ndarray, n_samp: int, transpose_axes: tuple @@ -70,15 +71,7 @@ def compute( con : np.ndarray wPLI connectivity matrix with shape (n_epoch, n_freq, 2*n_ch, 2*n_ch). """ - if self._backend == "metal": - return self._compute_metal(complex_signal, n_samp, transpose_axes) - elif self._backend == "cuda_kernel": - return self._compute_cuda(complex_signal, n_samp, transpose_axes) - elif self._backend == "torch": - return self._compute_torch(complex_signal, n_samp, transpose_axes) - elif self._backend == "numba": - return self._compute_numba(complex_signal, n_samp, transpose_axes) - return self._compute_numpy(complex_signal, n_samp, transpose_axes) + return super().compute(complex_signal, n_samp, transpose_axes) def _compute_metal( self, complex_signal: np.ndarray, n_samp: int, transpose_axes: tuple diff --git a/tests/test_sync.py b/tests/test_sync.py index 2e5ca5a..32669df 100644 --- a/tests/test_sync.py +++ b/tests/test_sync.py @@ -5,13 +5,14 @@ implementation to ensure numerical correctness. """ +import warnings from unittest.mock import patch import numpy as np import pytest from hypyp.analyses import compute_sync -from hypyp.sync import get_metric +from hypyp.sync import METRICS, get_metric from hypyp.sync.accorr import ACCorr from hypyp.sync.base import ( BaseMetric, @@ -25,6 +26,63 @@ from tests.accorr_reference import accorr_reference +#: Written by hand on purpose: the tests below must not derive their cases +#: from supports(), the function they check. +EXPECTED_BACKENDS = { + mode: {"numpy", "numba", "torch", "cuda_kernel"} + | ({"metal"} if mode in {"pli", "wpli", "accorr"} else set()) + for mode in ( + "plv", + "ccorr", + "accorr", + "coh", + "imcoh", + "pli", + "wpli", + "envcorr", + "powcorr", + ) +} +ALL_BACKENDS = ("numpy", "numba", "torch", "metal", "cuda_kernel") + +#: Also by hand: the routing test must not read the method names from +#: BaseMetric._BACKEND_METHODS, the table it checks. +EXPECTED_METHODS = { + "numpy": "_compute_numpy", + "numba": "_compute_numba", + "torch": "_compute_torch", + "metal": "_compute_metal", + "cuda_kernel": "_compute_cuda", +} + + +def spy_on_kernel(module_name, function_name): + """ + Patch a kernel function of ``hypyp.sync.kernels`` with a spy that still + runs it, to prove the kernel itself was entered. The metrics import their + kernel inside the method, so patching the module attribute is seen. + """ + import importlib + + module = importlib.import_module(f"hypyp.sync.kernels.{module_name}") + return patch.object( + module, function_name, side_effect=getattr(module, function_name) + ) + + +def spy_on(cls, method_name): + """ + Patch ``cls.`` with a spy that still runs the real method. + + Asserting ``metric._backend == 'metal'`` only proves what the metric + reports. The spy proves that ``compute`` really went through the method of + that backend, which is what a broken dispatch would get wrong. + """ + return patch.object( + cls, method_name, autospec=True, side_effect=getattr(cls, method_name) + ) + + class TestAccorrReference: """Basic properties of the reference implementation.""" @@ -225,20 +283,6 @@ def test_plv_symmetry(self, complex_signal): result[e, f], result[e, f].T, rtol=1e-10, atol=1e-12 ) - @pytest.mark.skipif(not METAL_AVAILABLE, reason="Metal not available") - def test_plv_metal_vs_numpy(self, complex_signal): - """Metal PLV should match numpy PLV within float32 tolerance.""" - from hypyp.sync.plv import PLV - - n_samp = complex_signal.shape[3] - result_np = PLV(optimization=None).compute( - complex_signal, n_samp, self.TRANSPOSE_AXES - ) - result_metal = PLV(optimization="metal").compute( - complex_signal, n_samp, self.TRANSPOSE_AXES - ) - np.testing.assert_allclose(result_metal, result_np, rtol=1e-5, atol=1e-5) - @pytest.mark.skipif(not CUPY_AVAILABLE, reason="CuPy not available") def test_plv_cuda_vs_numpy(self, complex_signal): """CUDA PLV should match numpy PLV exactly (both float64).""" @@ -356,23 +400,6 @@ def test_ccorr_torch_vs_numpy(self, complex_signal): else: np.testing.assert_allclose(result_torch, result_np, rtol=1e-9, atol=1e-10) - @pytest.mark.skipif(not METAL_AVAILABLE, reason="Metal not available") - def test_ccorr_metal_vs_numpy(self, complex_signal): - """Metal CCorr should match numpy CCorr within float32 tolerance. - - Uses Kahan summation with fastMath=OFF to preserve IEEE-754 compliance. - """ - from hypyp.sync.ccorr import CCorr - - n_samp = complex_signal.shape[3] - result_np = CCorr(optimization=None).compute( - complex_signal, n_samp, self.TRANSPOSE_AXES - ) - result_metal = CCorr(optimization="metal").compute( - complex_signal, n_samp, self.TRANSPOSE_AXES - ) - np.testing.assert_allclose(result_metal, result_np, rtol=1e-5, atol=1e-5) - @pytest.mark.skipif(not CUPY_AVAILABLE, reason="CuPy not available") def test_ccorr_cuda_vs_numpy(self, complex_signal): """CUDA CCorr should match numpy CCorr exactly (both float64).""" @@ -456,20 +483,6 @@ def test_coh_symmetry(self, complex_signal): result[e, f], result[e, f].T, rtol=1e-10, atol=1e-12 ) - @pytest.mark.skipif(not METAL_AVAILABLE, reason="Metal not available") - def test_coh_metal_vs_numpy(self, complex_signal): - """Metal Coh should match numpy Coh within float32 tolerance.""" - from hypyp.sync.coh import Coh - - n_samp = complex_signal.shape[3] - result_np = Coh(optimization=None).compute( - complex_signal, n_samp, self.TRANSPOSE_AXES - ) - result_metal = Coh(optimization="metal").compute( - complex_signal, n_samp, self.TRANSPOSE_AXES - ) - np.testing.assert_allclose(result_metal, result_np, rtol=1e-5, atol=1e-5) - @pytest.mark.skipif(not CUPY_AVAILABLE, reason="CuPy not available") def test_coh_cuda_vs_numpy(self, complex_signal): """CUDA Coh should match numpy Coh exactly (both float64).""" @@ -553,20 +566,6 @@ def test_imcoh_symmetry(self, complex_signal): result[e, f], result[e, f].T, rtol=1e-10, atol=1e-12 ) - @pytest.mark.skipif(not METAL_AVAILABLE, reason="Metal not available") - def test_imcoh_metal_vs_numpy(self, complex_signal): - """Metal ImCoh should match numpy ImCoh within float32 tolerance.""" - from hypyp.sync.imaginary_coh import ImCoh - - n_samp = complex_signal.shape[3] - result_np = ImCoh(optimization=None).compute( - complex_signal, n_samp, self.TRANSPOSE_AXES - ) - result_metal = ImCoh(optimization="metal").compute( - complex_signal, n_samp, self.TRANSPOSE_AXES - ) - np.testing.assert_allclose(result_metal, result_np, rtol=1e-5, atol=1e-5) - @pytest.mark.skipif(not CUPY_AVAILABLE, reason="CuPy not available") def test_imcoh_cuda_vs_numpy(self, complex_signal): """CUDA ImCoh should match numpy ImCoh exactly (both float64).""" @@ -755,9 +754,16 @@ def test_pli_metal_vs_numpy(self, complex_signal): result_np = PLI(optimization=None).compute( complex_signal, n_samp, self.TRANSPOSE_AXES ) - result_metal = PLI(optimization="metal").compute( - complex_signal, n_samp, self.TRANSPOSE_AXES - ) + metric_metal = PLI(optimization="metal") + # A silent fallback to numpy would make this comparison numpy-vs-numpy + # and therefore vacuous: check the backend, then that the Metal kernel + # function itself is entered. + assert metric_metal._backend == "metal" + with spy_on_kernel("metal_phase", "pli_metal") as spy: + result_metal = metric_metal.compute( + complex_signal, n_samp, self.TRANSPOSE_AXES + ) + assert spy.call_count == 1 # Float32 precision — sign() near zero can flip np.testing.assert_allclose(result_metal, result_np, rtol=1e-2, atol=1e-2) @@ -771,7 +777,11 @@ def test_pli_metal_large_channels(self): (2, 1, 256, 256) ) n_samp = sig.shape[3] - result = PLI(optimization="metal").compute(sig, n_samp, self.TRANSPOSE_AXES) + metric = PLI(optimization="metal") + assert metric._backend == "metal" + with spy_on_kernel("metal_phase", "pli_metal") as spy: + result = metric.compute(sig, n_samp, self.TRANSPOSE_AXES) + assert spy.call_count == 1 assert result.shape == (2, 1, 256, 256) assert not np.any(np.isnan(result)) assert np.allclose(np.diagonal(result[0, 0]), 0) # diagonal = 0 @@ -955,20 +965,6 @@ def test_envcorr_torch_vs_numpy(self, complex_signal): else: np.testing.assert_allclose(result_torch, result_np, rtol=1e-9, atol=1e-10) - @pytest.mark.skipif(not METAL_AVAILABLE, reason="Metal not available") - def test_envcorr_metal_vs_numpy(self, complex_signal): - """Metal EnvCorr should match numpy EnvCorr within float32 tolerance.""" - from hypyp.sync.envelope_corr import EnvCorr - - n_samp = complex_signal.shape[3] - result_np = EnvCorr(optimization=None).compute( - complex_signal, n_samp, self.TRANSPOSE_AXES - ) - result_metal = EnvCorr(optimization="metal").compute( - complex_signal, n_samp, self.TRANSPOSE_AXES - ) - np.testing.assert_allclose(result_metal, result_np, rtol=1e-5, atol=1e-5) - @pytest.mark.skipif(not CUPY_AVAILABLE, reason="CuPy not available") def test_envcorr_cuda_vs_numpy(self, complex_signal): """CUDA EnvCorr should match numpy EnvCorr exactly (both float64).""" @@ -1052,20 +1048,6 @@ def test_powcorr_torch_vs_numpy(self, complex_signal): else: np.testing.assert_allclose(result_torch, result_np, rtol=1e-9, atol=1e-10) - @pytest.mark.skipif(not METAL_AVAILABLE, reason="Metal not available") - def test_powcorr_metal_vs_numpy(self, complex_signal): - """Metal PowCorr should match numpy PowCorr within float32 tolerance.""" - from hypyp.sync.pow_corr import PowCorr - - n_samp = complex_signal.shape[3] - result_np = PowCorr(optimization=None).compute( - complex_signal, n_samp, self.TRANSPOSE_AXES - ) - result_metal = PowCorr(optimization="metal").compute( - complex_signal, n_samp, self.TRANSPOSE_AXES - ) - np.testing.assert_allclose(result_metal, result_np, rtol=1e-5, atol=1e-5) - @pytest.mark.skipif(not CUPY_AVAILABLE, reason="CuPy not available") def test_powcorr_cuda_vs_numpy(self, complex_signal): """CUDA PowCorr should match numpy PowCorr exactly (both float64).""" @@ -1089,9 +1071,13 @@ def test_wpli_metal_vs_numpy(self, complex_signal): result_np = WPLI(optimization=None).compute( complex_signal, n_samp, self.TRANSPOSE_AXES ) - result_metal = WPLI(optimization="metal").compute( - complex_signal, n_samp, self.TRANSPOSE_AXES - ) + metric_metal = WPLI(optimization="metal") + assert metric_metal._backend == "metal" + with spy_on_kernel("metal_phase", "wpli_metal") as spy: + result_metal = metric_metal.compute( + complex_signal, n_samp, self.TRANSPOSE_AXES + ) + assert spy.call_count == 1 np.testing.assert_allclose(result_metal, result_np, rtol=1e-2, atol=1e-2) @pytest.mark.skipif(not CUPY_AVAILABLE, reason="CuPy not available") @@ -1123,9 +1109,13 @@ def test_accorr_metal_vs_numpy(self, complex_signal): result_np = ACCorr(optimization=None, show_progress=False).compute( complex_signal, n_samp, self.TRANSPOSE_AXES ) - result_metal = ACCorr(optimization="metal", show_progress=False).compute( - complex_signal, n_samp, self.TRANSPOSE_AXES - ) + metric_metal = ACCorr(optimization="metal", show_progress=False) + assert metric_metal._backend == "metal" + with spy_on_kernel("metal_accorr", "accorr_metal") as spy: + result_metal = metric_metal.compute( + complex_signal, n_samp, self.TRANSPOSE_AXES + ) + assert spy.call_count == 1 np.testing.assert_allclose(result_metal, result_np, rtol=1e-5, atol=1e-5) @pytest.mark.skipif(not CUPY_AVAILABLE, reason="CuPy not available") @@ -1237,3 +1227,574 @@ def test_priority_parameter_propagated_via_get_metric(self): """get_metric passes priority through to the metric class.""" m = get_metric("accorr", optimization="auto", priority=["numba"]) assert m._priority == ["numba"] + + +class TestBackendCapability: + """ + A requested backend must either run, or degrade with a warning. + + Only PLI, wPLI and ACCorr have Metal kernels — torch/MPS is the intended + GPU path for the six einsum metrics (see hypyp/sync/base.py AUTO_PRIORITY + rationale and the support matrix in hypyp/sync/README.md). Requesting + 'metal' for a metric that has no Metal kernel must therefore be reported, + never silently answered with numpy. + + These tests patch the availability flags instead of gating on hardware, so + the capability contract is verified on any machine including CI. + """ + + METAL_CAPABLE = {"pli", "wpli", "accorr"} + + def test_capability_matrix(self): + """supports() must answer exactly the hand-written support matrix.""" + assert set(METRICS) == set(EXPECTED_BACKENDS) + for mode, cls in METRICS.items(): + actual = {b for b in ALL_BACKENDS if cls.supports(b)} + assert actual == EXPECTED_BACKENDS[mode], mode + assert cls.supports("not_a_backend") is False + + def test_supports_reflects_the_implemented_methods(self): + """supports() must be derived from the code, not a hand-kept list.""" + for mode, cls in METRICS.items(): + assert cls.supports("metal") == hasattr(cls, "_compute_metal") + assert cls.supports("numpy") is True + assert cls.supports("numba") == hasattr(cls, "_compute_numba") + + def test_metal_capability_matches_documented_support_matrix(self): + """Exactly PLI/wPLI/ACCorr expose a Metal kernel.""" + actual = {mode for mode, cls in METRICS.items() if cls.supports("metal")} + assert actual == self.METAL_CAPABLE + + @pytest.mark.parametrize("mode", sorted(METRICS)) + def test_metal_request_never_silently_degrades(self, mode, complex_signal): + """ + optimization='metal' either resolves to metal, or warns and uses numpy. + + Regression test: before the capability check, _resolve_optimization + granted ('metal', 'mps') to every metric, and compute() then fell + through its if/elif chain into _compute_numpy — so six metrics returned + a numpy result while reporting _backend == 'metal', with no warning. + """ + with patch("hypyp.sync.base.METAL_AVAILABLE", True): + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + metric = get_metric(mode, optimization="metal") + + if mode in self.METAL_CAPABLE: + assert metric._backend == "metal" + else: + assert metric._backend == "numpy", ( + f"{mode}: asked for metal, resolved to {metric._backend!r} " + f"but has no Metal kernel" + ) + messages = [ + str(w.message) for w in caught if issubclass(w.category, UserWarning) + ] + assert any("Metal" in m for m in messages), ( + f"{mode}: degraded to numpy without warning (messages: {messages})" + ) + # The fallback must also compute: same result as a plain numpy + # metric, through the numpy method. + n_samp = complex_signal.shape[3] + axes = (0, 1, 3, 2) + expected = get_metric(mode).compute(complex_signal, n_samp, axes) + with spy_on(type(metric), "_compute_numpy") as spy: + result = metric.compute(complex_signal, n_samp, axes) + assert spy.call_count == 1 + np.testing.assert_array_equal(result, expected) + + @pytest.mark.parametrize( + "mode, backend", + [ + (mode, backend) + for mode in sorted(EXPECTED_BACKENDS) + for backend in sorted(EXPECTED_BACKENDS[mode]) + ], + ) + def test_compute_routes_to_the_method_of_the_backend(self, mode, backend): + """ + compute() must call the ``_compute_*`` method of the resolved backend, + with its arguments, and return its result, for every backend each + metric implements. No hardware is needed: the method is replaced by a + stub. + """ + cls = METRICS[mode] + method_name = EXPECTED_METHODS[backend] + metric = cls() + metric._backend = backend + sentinel = object() + with patch.object(cls, method_name, autospec=True) as stub: + stub.return_value = sentinel + result = metric.compute("signal", 7, (0, 1, 3, 2)) + assert result is sentinel + stub.assert_called_once_with(metric, "signal", 7, (0, 1, 3, 2)) + + @pytest.mark.parametrize("mode", sorted(set(METRICS) - {"pli", "wpli", "accorr"})) + @pytest.mark.parametrize("priority", [["metal", "torch"], ["metal"]]) + def test_auto_priority_unsupported_backend_stays_on_numpy(self, mode, priority): + """ + A priority list that reaches an available backend the metric cannot + run keeps computing in numpy, as before, but now says so. + + The 0.6 series changes no computed value: moving on to the next + backend of the list (torch, or the numba fallback) would. + """ + with ( + patch("hypyp.sync.base.METAL_AVAILABLE", True), + patch("hypyp.sync.base.TORCH_AVAILABLE", True), + patch("hypyp.sync.base.MPS_AVAILABLE", True), + patch("hypyp.sync.base.NUMBA_AVAILABLE", True), + ): + with pytest.warns(UserWarning, match="no Metal implementation"): + metric = get_metric(mode, optimization="auto", priority=priority) + assert (metric._backend, metric._device) == ("numpy", "cpu"), ( + f"{mode}: priority={priority} resolved to {metric._backend!r}" + ) + + def test_unknown_backend_fails_closed(self, complex_signal): + """ + An unrecognised _backend must raise, not quietly compute in numpy. + + This is the fail-closed guarantee: the original if/elif chains ended in + a bare `return self._compute_numpy(...)`, so any unhandled backend value + became indistinguishable from the default. + """ + from hypyp.sync.plv import PLV + + n_samp = complex_signal.shape[3] + metric = PLV() + metric._backend = "not_a_backend" + with pytest.raises(ValueError) as excinfo: + metric.compute(complex_signal, n_samp, (0, 1, 3, 2)) + # The error must be usable: it names the metric, the offending + # backend and the backends that do exist for this metric. + message = str(excinfo.value) + assert "'plv'" in message + assert "'not_a_backend'" in message + assert "numpy" in message + + def test_unimplemented_backend_computes_in_numpy_with_a_warning( + self, complex_signal + ): + """ + A known backend the metric does not implement computes in numpy, as + before, but says so: the 0.6 series changes no computed value. + """ + from hypyp.sync.plv import PLV + + n_samp = complex_signal.shape[3] + axes = (0, 1, 3, 2) + metric = PLV() + metric._backend = "metal" + with pytest.warns(UserWarning, match="'plv' has no Metal implementation"): + result = metric.compute(complex_signal, n_samp, axes) + assert isinstance(result, np.ndarray) + np.testing.assert_array_equal( + result, PLV()._compute_numpy(complex_signal, n_samp, axes) + ) + + def test_priority_fallback_warning_names_the_skipped_backend(self): + """ + When the only backend of a priority list has no implementation and + cannot run on the machine either, the fallback warning must give the + first reason. + + Before, priority=['metal'] on an einsum metric warned "No GPU backend + available" on a CUDA machine, where a GPU backend was available, + without mentioning that the metric has no Metal kernel. + """ + with ( + patch("hypyp.sync.base.METAL_AVAILABLE", False), + patch("hypyp.sync.base.TORCH_AVAILABLE", True), + patch("hypyp.sync.base.MPS_AVAILABLE", False), + patch("hypyp.sync.base.CUDA_AVAILABLE", True), + ): + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + metric = get_metric("plv", optimization="auto", priority=["metal"]) + + assert metric._backend in ("numba", "numpy") + messages = [ + str(w.message) for w in caught if issubclass(w.category, UserWarning) + ] + assert any("'plv' has no Metal implementation" in m for m in messages), ( + f"fallback warning does not explain the skip (messages: {messages})" + ) + assert not any("No GPU backend available" in m for m in messages) + # The message must stay true when the CPU fallback is numba: it speaks + # of GPU backends only. + assert any("no GPU backend of the priority list" in m for m in messages) + + @staticmethod + def _legacy_metric(with_helper=False): + """A third-party metric written against the pre-0.6.2 contract: it + overrides ``compute``, branches on ``self._backend`` itself and has no + ``_compute_*`` method.""" + from hypyp.sync.base import BaseMetric + + class LegacyMetric(BaseMetric): + name = "legacy" + + def compute(self, complex_signal, n_samp, transpose_axes): + base_result = super().compute(complex_signal, n_samp, transpose_axes) + return self._backend, base_result + + class LegacyWithHelper(LegacyMetric): + # Same contract, but the author happened to name a helper like the + # methods of the current contract. Still its own dispatch. + def _compute_numpy(self, complex_signal, n_samp, transpose_axes): + return "helper" + + return LegacyWithHelper if with_helper else LegacyMetric + + @pytest.mark.parametrize( + "kwargs, expected", + [ + (dict(optimization=None), "numpy"), + (dict(optimization="numba"), "numba"), + (dict(optimization="torch"), "torch"), + (dict(optimization="metal"), "metal"), + (dict(optimization="auto", priority=["metal", "torch"]), "metal"), + (dict(optimization="auto", priority=["torch"]), "torch"), + ], + ) + @pytest.mark.parametrize("with_helper", [False, True]) + def test_legacy_subclass_keeps_its_own_dispatch( + self, kwargs, expected, with_helper + ): + """ + A subclass of the earlier contract is granted the backend it asks for, + as before the capability check, and without a "no implementation" + warning: the base class cannot see inside its ``compute``. + """ + legacy_cls = self._legacy_metric(with_helper) + with ( + patch("hypyp.sync.base.METAL_AVAILABLE", True), + patch("hypyp.sync.base.TORCH_AVAILABLE", True), + patch("hypyp.sync.base.MPS_AVAILABLE", True), + patch("hypyp.sync.base.NUMBA_AVAILABLE", True), + ): + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + metric = legacy_cls(**kwargs) + assert metric._backend == expected + assert not any("implementation" in str(w.message) for w in caught) + # super().compute() used to be an abstract method with an empty body: + # it returned None and must still do so, not raise. + assert metric.compute(None, 0, None) == (expected, None) + + @pytest.mark.parametrize("gpu", [True, False]) + def test_numpy_only_metric_never_gets_numba(self, gpu, complex_signal): + """ + A metric of the current contract that implements numpy alone must + resolve to numpy under 'auto', even when numba is installed: the CPU + fallback has to respect capability like every other path. + """ + from hypyp.sync.base import BaseMetric + + class NumpyOnly(BaseMetric): + name = "numpy_only" + + def _compute_numpy(self, complex_signal, n_samp, transpose_axes): + return np.ones(1) + + with ( + patch("hypyp.sync.base.NUMBA_AVAILABLE", True), + patch("hypyp.sync.base.TORCH_AVAILABLE", gpu), + patch("hypyp.sync.base.MPS_AVAILABLE", gpu), + patch("hypyp.sync.base.CUDA_AVAILABLE", False), + ): + with warnings.catch_warnings(): + warnings.simplefilter("ignore") + metric = NumpyOnly(optimization="auto") + assert metric._backend == "numpy" + n_samp = complex_signal.shape[3] + assert metric.compute(complex_signal, n_samp, (0, 1, 3, 2)).shape == (1,) + + def test_subclass_of_builtin_with_its_own_backend(self, complex_signal): + """ + A third-party subclass of a built-in metric that handles a backend in + its own ``compute`` and delegates the rest to its parent keeps working: + the backend is granted without warning, its own branch runs, and the + delegation still computes. + """ + from hypyp.sync.plv import PLV + + class CustomPLV(PLV): + def compute(self, complex_signal, n_samp, transpose_axes): + if self._backend == "metal": + return "custom Metal" + return super().compute(complex_signal, n_samp, transpose_axes) + + with patch("hypyp.sync.base.METAL_AVAILABLE", True): + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + metric = CustomPLV(optimization="metal") + assert metric._backend == "metal" + assert not caught + assert metric.compute(None, 0, None) == "custom Metal" + + n_samp = complex_signal.shape[3] + axes = (0, 1, 3, 2) + delegated = CustomPLV().compute(complex_signal, n_samp, axes) + # Compared with the numpy method itself, and checked to be an array: + # two None results would otherwise compare equal. + assert isinstance(delegated, np.ndarray) + np.testing.assert_array_equal( + delegated, PLV()._compute_numpy(complex_signal, n_samp, axes) + ) + # The built-in parent itself stays capability-checked. + assert PLV.supports("metal") is False + + def test_dispatch_flag_set_by_a_mixin(self): + """ + The class that sets ``_dispatch_via_table`` need not define + ``compute``: a mixin can carry the flag. The metric is then + capability-checked like its built-in parent, and does not crash. + """ + from hypyp.sync.plv import PLV + + class DispatchPolicy: + _dispatch_via_table = True + + class MixedPLV(DispatchPolicy, PLV): + pass + + assert MixedPLV.supports("numpy") is True + assert MixedPLV.supports("metal") is False + with patch("hypyp.sync.base.METAL_AVAILABLE", True): + with pytest.warns(UserWarning, match="no Metal implementation"): + metric = MixedPLV(optimization="metal") + assert metric._backend == "numpy" + + def test_rebinding_the_parent_compute_keeps_the_capability_check(self): + """ + ``compute = PLV.compute`` in a subclass is the same function as the + one the flag of PLV vouches for, not a dispatch of its own: the + subclass is still capability-checked. + """ + from hypyp.sync.plv import PLV + + class Alias(PLV): + compute = PLV.compute + + assert Alias.supports("numpy") is True + assert Alias.supports("metal") is False + with patch("hypyp.sync.base.METAL_AVAILABLE", True): + with pytest.warns(UserWarning, match="no Metal implementation"): + metric = Alias(optimization="metal") + assert metric._backend == "numpy" + + def test_auto_keeps_numba_handled_inside_compute(self, complex_signal): + """ + A descendant of a built-in metric that hides ``_compute_numba`` and + handles numba inside its own ``compute`` still gets numba from the + CPU fallback of ``'auto'``, as before the capability check, and its + own branch runs. + """ + from hypyp.sync.plv import PLV + + class OwnNumba(PLV): + _compute_numba = None + + def compute(self, complex_signal, n_samp, transpose_axes): + if self._backend == "numba": + return "own numba" + return super().compute(complex_signal, n_samp, transpose_axes) + + n_samp = complex_signal.shape[3] + axes = (0, 1, 3, 2) + with ( + patch("hypyp.sync.base.NUMBA_AVAILABLE", True), + patch("hypyp.sync.base.TORCH_AVAILABLE", False), + patch("hypyp.sync.base.MPS_AVAILABLE", False), + patch("hypyp.sync.base.CUDA_AVAILABLE", False), + ): + with warnings.catch_warnings(): + warnings.simplefilter("ignore") + automatic = OwnNumba(optimization="auto") + requested = OwnNumba(optimization="numba") + assert automatic._backend == "numba" + assert automatic.compute(complex_signal, n_samp, axes) == "own numba" + assert requested._backend == "numba" + assert requested.compute(complex_signal, n_samp, axes) == "own numba" + + @pytest.mark.parametrize("delegates", [False, True]) + def test_backend_method_added_by_a_subclass_needs_the_flag( + self, complex_signal, delegates + ): + """ + Before the table dispatch, a ``_compute_metal`` added to a subclass of + PLV was never called: the request computed in numpy. That result is + kept, with a warning, whether the subclass inherits ``compute`` or + overrides it only to delegate. The added method is used once the + subclass sets ``_dispatch_via_table`` itself. + """ + from hypyp.sync.plv import PLV + + class AddsMetal(PLV): + def _compute_metal(self, complex_signal, n_samp, transpose_axes): + return "added Metal" + + def _compute_numpy(self, complex_signal, n_samp, transpose_axes): + # An overridden method of the parent is still honoured. + return 2 * super()._compute_numpy( + complex_signal, n_samp, transpose_axes + ) + + if delegates: + + class AddsMetal(AddsMetal): # noqa: F811 + def compute(self, complex_signal, n_samp, transpose_axes): + return super().compute(complex_signal, n_samp, transpose_axes) + + class Migrated(AddsMetal): + _dispatch_via_table = True + + n_samp = complex_signal.shape[3] + axes = (0, 1, 3, 2) + expected = 2 * PLV()._compute_numpy(complex_signal, n_samp, axes) + with patch("hypyp.sync.base.METAL_AVAILABLE", True): + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + metric = AddsMetal(optimization="metal") + result = metric.compute(complex_signal, n_samp, axes) + migrated = Migrated(optimization="metal") + assert any("no Metal implementation" in str(w.message) for w in caught) + assert isinstance(result, np.ndarray) + np.testing.assert_array_equal(result, expected) + assert Migrated.supports("metal") is True + assert migrated._backend == "metal" + assert migrated.compute(complex_signal, n_samp, axes) == "added Metal" + + def test_dispatch_flag_set_by_a_mixin_listed_after_the_base(self): + """ + A mixin that carries the flag can come after ``BaseMetric`` in the + bases, where no class defines ``compute`` any more. A metric with its + own ``compute`` is then trusted with the backend it requests, and one + without is capability-checked; neither crashes. + """ + + class DispatchPolicy: + _dispatch_via_table = True + + class OwnDispatch(BaseMetric, DispatchPolicy): + name = "own_dispatch" + + def compute(self, complex_signal, n_samp, transpose_axes): + return self._backend + + class TableOnly(BaseMetric, DispatchPolicy): + name = "table_only" + + def _compute_numpy(self, complex_signal, n_samp, transpose_axes): + return "numpy" + + with patch("hypyp.sync.base.NUMBA_AVAILABLE", True): + assert OwnDispatch.supports("numba") is True + metric = OwnDispatch(optimization="numba") + assert metric.compute(None, 0, None) == "numba" + assert TableOnly.supports("numpy") is True + assert TableOnly.supports("numba") is False + + def test_delegating_descendant_of_numpy_only_metric_still_computes(self): + """ + A descendant that overrides ``compute`` only to delegate is trusted + with numba by the automatic CPU fallback, like any class with its own + ``compute``. No numba method exists, so the dispatch computes in numpy + and warns instead of failing. + """ + from hypyp.sync.base import BaseMetric + + class NumpyOnly(BaseMetric): + name = "numpy_only" + _dispatch_via_table = True + + def _compute_numpy(self, complex_signal, n_samp, transpose_axes): + return "numpy result" + + class Delegating(NumpyOnly): + def compute(self, complex_signal, n_samp, transpose_axes): + return super().compute(complex_signal, n_samp, transpose_axes) + + with ( + patch("hypyp.sync.base.NUMBA_AVAILABLE", True), + patch("hypyp.sync.base.TORCH_AVAILABLE", False), + patch("hypyp.sync.base.MPS_AVAILABLE", False), + patch("hypyp.sync.base.CUDA_AVAILABLE", False), + ): + with warnings.catch_warnings(): + warnings.simplefilter("ignore") + metric = Delegating(optimization="auto") + assert metric._backend == "numba" + with pytest.warns(UserWarning, match="no numba implementation"): + assert metric.compute(None, 0, None) == "numpy result" + + def test_classmethod_implementation_is_recognised(self): + """A ``_compute_*`` method declared as a classmethod is an + implementation like any other.""" + from hypyp.sync.base import BaseMetric + + class ClassLevel(BaseMetric): + name = "class_level" + _dispatch_via_table = True + + @classmethod + def _compute_numpy(cls, complex_signal, n_samp, transpose_axes): + return "numpy result" + + assert ClassLevel.supports("numpy") is True + assert ClassLevel().compute(None, 0, None) == "numpy result" + + def test_supports_ignores_placeholders(self): + """supports() must not count a non-callable attribute, nor the default + ``_compute_numpy`` of the base class, as an implementation.""" + from hypyp.sync.base import BaseMetric + + class Placeholder(BaseMetric): + name = "placeholder" + _compute_torch = None + + def _compute_numpy(self, complex_signal, n_samp, transpose_axes): + return np.ones(1) + + class EmptyMetric(BaseMetric): + name = "empty" + + assert Placeholder.supports("numpy") is True + assert Placeholder.supports("torch") is False + assert EmptyMetric.supports("numpy") is False + assert BaseMetric.supports("numpy") is False + + def test_torch_only_metric_reports_its_backends(self): + """A metric of the current contract without numpy runs the backend it + has and names what is missing otherwise.""" + from hypyp.sync.base import BaseMetric + + class TorchOnly(BaseMetric): + name = "torch_only" + + def _compute_torch(self, complex_signal, n_samp, transpose_axes): + return "torch result" + + metric = TorchOnly() + with pytest.raises(NotImplementedError, match="_compute_numpy"): + metric.compute(None, 0, None) + metric._backend = "torch" + assert metric.compute(None, 0, None) == "torch result" + metric._backend = "metal" + with pytest.raises( + ValueError, match=r"implemented for this metric: \['torch'\]" + ): + metric.compute(None, 0, None) + + def test_missing_numpy_implementation_is_reported(self, complex_signal): + """A metric with neither ``compute`` nor ``_compute_numpy`` says so.""" + from hypyp.sync.base import BaseMetric + + class EmptyMetric(BaseMetric): + name = "empty" + + n_samp = complex_signal.shape[3] + with pytest.raises(NotImplementedError, match="_compute_numpy"): + EmptyMetric().compute(complex_signal, n_samp, (0, 1, 3, 2))