Source code for typed_lisa_toolkit.shop.transforms

"""Functions for transforming data and waveforms between time-domain, frequency-domain, and time-frequency plane."""  # noqa: E501

import os
import warnings
from types import ModuleType
from typing import Literal, cast, overload

from .. import (
    _constructors,  # pyright: ignore[reportPrivateUsage]
    utils,
)
from ..types import AnyArray, AnyAxis, Axis, AxLike, Grid2DCartesian, Linspace, data
from ..types import representations as reps


def _import_wdm_transform() -> ModuleType:
    try:
        import wdm_transform
    except ImportError:
        msg = (
            "The `wdm_transform` package is required for WDM transformations."
            "Install it with: pip install wdm-transform"
        )
        raise ImportError(
            msg,
        ) from None
    else:
        return wdm_transform


def _conventionaize(ary: "AnyArray") -> "AnyArray":
    return ary[:, None, None, None, ...]


# Meyer window default parameters
DEFAULT_WINDOW_A = 1 / 3
DEFAULT_WINDOW_D = 1.0


@overload
def time2freq(
    tsd: data.TSData,
    /,
    *,
    keep_time: Literal[True] = True,
) -> data.TimedFSData: ...
@overload
def time2freq(
    tsd: data.TSData,
    /,
    *,
    keep_time: Literal[False],
) -> data.FSData: ...
@overload
def time2freq(
    ts: reps.TimeSeries[Axis[Linspace]],
    /,
    *,
    keep_time: bool = True,
) -> reps.UniformFrequencySeries: ...
[docs] def time2freq( td: reps.TimeSeries[Axis[Linspace]] | data.TSData, /, *, keep_time: bool = True, ): """Convert time-domain representation or data to frequency-domain representation or data using the real FFT. .. note:: When the input is a :class:`~types.TimeSeries` instance, the ``keep_time`` argument is ignored. """ # noqa: E501 xp = td.xp try: fft = xp.fft except AttributeError as e: msg = ( f"{xp.__name__} does not support FFT operations, " "which are required for `time2freq`." ) raise NotImplementedError(msg) from e n_times = len(td.times) last_freq = (n_times // 2) / n_times / td.times.ax.step n_freq = n_times // 2 + 1 freqs = _constructors.axis(_constructors.linspace(0.0, last_freq, n_freq)) signal = cast("AnyArray", fft.rfft(td.get_kernel() * td.times.ax.step, axis=-1)) if isinstance(td, reps.TimeSeries): return _constructors.frequency_series( frequencies=freqs, entries=signal, ) fsd = _constructors.fsdata( frequencies=freqs, entries=signal, channels=td.channel_names, name=td.name, ) if keep_time: fsd = fsd.set_times(td.times) return fsd
@overload def freq2time(fsd: data.FSData, /, *, times: AnyAxis | AxLike) -> data.TSData: ... @overload def freq2time( fsd: reps.FrequencySeries[Axis[Linspace]], /, *, times: AnyAxis | AxLike, ) -> reps.UniformTimeSeries: ...
[docs] def freq2time( fd: reps.FrequencySeries[Axis[Linspace]] | data.FSData, /, *, times: AnyAxis | AxLike, ): """Convert frequency-domain representation or data to time-domain representation or data using the inverse real FFT.""" # noqa: E501 xp = fd.xp try: fft = xp.fft except AttributeError as e: msg = ( f"{xp.__name__} does not support FFT operations, " "which are required for `freq2time`." ) raise NotImplementedError(msg) from e _times = ( Linspace.make(times) if not isinstance(times, Axis) else Linspace.make(times.ax) ) is_even = len(fd.frequencies) % 2 == 0 nyquist_freq = ( fd.frequencies.stop if is_even else fd.frequencies.stop + fd.frequencies.ax.step / 2 ) nyquist_dt = 1.0 / (2 * nyquist_freq) if _times.step < nyquist_dt and not xp.isclose( xp.asarray(_times.step), xp.asarray(nyquist_dt) ): warnings.warn("The time grid is denser than the Nyquist limit.", stacklevel=2) signal = cast( "AnyArray", fft.irfft(fd.get_kernel() / _times.step, n=len(_times), axis=-1) ) if isinstance(fd, reps.FrequencySeries): return _constructors.time_series(times=_times, entries=signal) return _constructors.tsdata( times=_times, entries=signal, channels=fd.channel_names, name=fd.name, )
def _wdm_backend_env(xp: "ModuleType") -> dict[str, str]: if xp.__name__ in ("numpy", "array_api_compat.numpy"): pass elif "jax" in xp.__name__: _backend = os.environ.get("WDM_BACKEND") if _backend is None: return {"WDM_BACKEND": "jax"} if _backend != "jax": msg = ( f"WDM_BACKEND is set to {_backend!r}, " "but the input data uses JAX arrays. " "Please set WDM_BACKEND to 'jax' or unset it to use the JAX backend." ) raise ValueError(msg) return {} @overload def time2wdm( tdata: data.TSData, /, *, Nt: int, # noqa: N803 Nf: int, # noqa: N803 ) -> data.WDMData[Grid2DCartesian[Axis[Linspace], Axis[Linspace]]]: ... @overload def time2wdm( tseries: reps.TimeSeries[Axis[Linspace]], /, *, Nt: int, # noqa: N803 Nf: int, # noqa: N803 ) -> reps.WDM[Grid2DCartesian[Axis[Linspace], Axis[Linspace]]]: ...
[docs] def time2wdm( tthing: data.TSData | reps.TimeSeries[Axis[Linspace]], /, *, Nt: int, # noqa: N803 Nf: int, # noqa: N803 ): """Transform a time series to WDM. .. note:: The convention for signal duration in :class:`WDM` is Nt*ΔT = N*Δt. This is not equivalent to the convention ``grid[-1] - grid[0]``, more useful for nonuniform grids, that you may be assuming elsewhere. Parameters ---------- tseries : TimeSeries Regularly-sampled time series of length at least Nf*Nt. Nt : int Length of WDM time grid. Nf : int Nf+1 is the length of the WDM frequency grid. Returns ------- WDM Transform of the first Nf*Nt points of the time series. Raises ------ ValueError If the time series is too small for the chosen values of Nf and Nt. """ _import_wdm_transform() from wdm_transform.transforms import ( forward_wdm as _forward_wdm, ) if isinstance(tthing, data.TSData): return _constructors.wdmdata( {key: time2wdm(val, Nt=Nt, Nf=Nf) for (key, val) in tthing.items()}, name=tthing.name, ) assert isinstance(tthing, reps.TimeSeries) # noqa: S101 tseries = tthing if Nt * Nf > tseries.times.ax.num: msg = "Time series too small for given Nf and Nt" raise ValueError(msg) tseries = tseries[: Nt * Nf] n_batches = tseries.entries.shape[0] if tseries.entries.shape[1:4] != (1, 1, 1): msg = "Currently only single-channel time series are supported by time2wdm." raise ValueError(msg) _entries = tseries.entries[:, 0, 0, 0] with utils.set_env(**_wdm_backend_env(tseries.xp)): coeffs = _forward_wdm( _entries, nt=Nt, nf=Nf, a=DEFAULT_WINDOW_A, d=DEFAULT_WINDOW_D, dt=tseries.times.ax.step, ).swapaxes(-2, -1) expected_shape = (n_batches, Nf + 1, Nt) if coeffs.shape != expected_shape: msg = ( "Unexpected shape of WDM coefficients: " f"expected {expected_shape}, got {coeffs.shape}" ) raise ValueError(msg) dT = Nf * tseries.times.ax.step # noqa: N806 dF = 0.5 / dT # noqa: N806 tgrid = Linspace(start=tseries.times.start, step=dT, num=Nt) fgrid = Linspace(start=0, step=dF, num=Nf + 1) return _constructors.wdm( times=tgrid, frequencies=fgrid, entries=_conventionaize(coeffs), )
@overload def wdm2time( wdmdata: data.WDMData[Grid2DCartesian[Axis[Linspace], Axis[Linspace]]], /, ) -> data.TSData: ... @overload def wdm2time( wdm: reps.WDM[Grid2DCartesian[Axis[Linspace], Axis[Linspace]]], /, ) -> reps.UniformTimeSeries: ...
[docs] def wdm2time( wdmthing: reps.WDM[Grid2DCartesian[Axis[Linspace], Axis[Linspace]]] | data.WDMData[Grid2DCartesian[Axis[Linspace], Axis[Linspace]]], /, ): """Transform WDM expansion to equivalent time series. Parameters ---------- wdm : WDM or WDMData WDM expansion with grid parameters Nf and Nt. Returns ------- TimeSeries or TSData Time series of size Nt*Nf. """ _import_wdm_transform() from wdm_transform.transforms import ( inverse_wdm as _inverse_wdm, ) if isinstance(wdmthing, data.WDMData): return _constructors.tsdata( {key: wdm2time(val) for (key, val) in wdmthing.items()}, name=wdmthing.name, ) assert isinstance(wdmthing, reps.WDM) # noqa: S101 wdm = wdmthing if wdm.entries.shape[1:4] != (1, 1, 1): msg = "Currently only single-channel WDMs are supported by wdm2time." raise ValueError(msg) # swapaxes are not specified in array-api so we have to ignore the type checks here _coeffs = wdm.entries[:, 0, 0, 0].swapaxes(-2, -1) # pyright: ignore[reportUnknownVariableType,reportUnknownMemberType,reportAttributeAccessIssue] entries = _inverse_wdm( coeffs=_coeffs, dt=wdm.dt, a=DEFAULT_WINDOW_A, d=DEFAULT_WINDOW_D, ) tgrid = Linspace(wdm.times.start, wdm.dt, wdm.ND) return _constructors.time_series(tgrid, _conventionaize(entries))
@overload def freq2wdm( fseries: reps.FrequencySeries[Axis[Linspace]], /, *, Nt: int, # noqa: N803 Nf: int, # noqa: N803 t0: float = 0.0, ) -> reps.WDM[Grid2DCartesian[Axis[Linspace], Axis[Linspace]]]: ... @overload def freq2wdm( fdata: data.FSData, /, *, Nt: int, # noqa: N803 Nf: int, # noqa: N803 t0: float = 0.0, ) -> data.WDMData[Grid2DCartesian[Axis[Linspace], Axis[Linspace]]]: ...
[docs] def freq2wdm( fthing: reps.FrequencySeries[Axis[Linspace]] | data.FSData, /, *, Nt: int, # noqa: N803 Nf: int, # noqa: N803 t0: float = 0.0, ): """Transform frequency series to WDM. Parameters ---------- fseries : :class:`~types.FrequencySeries` or :class:`~types.FSData` frequency series, assumed to be full-grid (from DC to Nyquist). Nt : int Number of WDM time bins. Nf : int WDM frequency grid has length Nf+1. t0 : float, optional Initial time of WDM time grid, by default 0.0 """ _import_wdm_transform() from wdm_transform.transforms import ( get_backend as _get_backend, ) if isinstance(fthing, data.FSData): return _constructors.wdmdata( {key: freq2wdm(val, Nt=Nt, Nf=Nf, t0=t0) for (key, val) in fthing.items()}, name=fthing.name, ) assert isinstance(fthing, reps.FrequencySeries), ( # noqa: S101 f"Expected a FrequencySeries input, got {type(fthing)}" ) fseries = fthing with utils.set_env(**_wdm_backend_env(fseries.xp)): backend = _get_backend() tseries_entries = backend.fft.irfft(fseries.entries, n=Nf * Nt) duration = 1 / fseries.frequencies.ax.step dt = duration / (Nf * Nt) tseries_grid = Linspace(start=t0, step=dt, num=Nf * Nt) tseries = _constructors.time_series(tseries_grid, tseries_entries) return time2wdm(tseries, Nt=Nt, Nf=Nf)
@overload def wdm2freq( wdmdata: data.WDMData[Grid2DCartesian[Axis[Linspace], Axis[Linspace]]], /, ) -> data.FSData: ... @overload def wdm2freq( wdm: reps.WDM[Grid2DCartesian[Axis[Linspace], Axis[Linspace]]], /, ) -> reps.UniformFrequencySeries: ...
[docs] def wdm2freq( wdmthing: reps.WDM[Grid2DCartesian[Axis[Linspace], Axis[Linspace]]] | data.WDMData[Grid2DCartesian[Axis[Linspace], Axis[Linspace]]], /, ): """Transform WDM expansion to a frequency series. .. note:: The input WDM representation is assumed to contain all frequencies from DC to Nyquist. """ _import_wdm_transform() from wdm_transform.transforms import ( frequency_wdm as _frequency_wdm, ) if isinstance(wdmthing, data.WDMData): return _constructors.fsdata( {key: wdm2freq(val) for (key, val) in wdmthing.items()}, name=wdmthing.name, ) assert isinstance(wdmthing, reps.WDM) # noqa: S101 wdm = wdmthing if wdm.entries.shape[1:4] != (1, 1, 1): msg = "Currently only single-channel WDMs are supported by wdm2freq." raise ValueError(msg) # swapaxes are not specified in array-api so we have to ignore the type checks here _coeffs = wdm.entries[:, 0, 0, 0].swapaxes(-2, -1) # pyright: ignore[reportUnknownVariableType,reportUnknownMemberType,reportAttributeAccessIssue] with utils.set_env(**_wdm_backend_env(wdm.xp)): wtfs = _frequency_wdm( _coeffs, dt=wdm.dt, a=DEFAULT_WINDOW_A, d=DEFAULT_WINDOW_D ) # wtfs is on a grid from fftfreq but we want rfftfreq _n = wtfs.shape[1] _num = _n // 2 + 1 if _n % 2 == 0 else (_n + 1) // 2 freqs = _constructors.linspace_from_step(start=0.0, step=wdm.df, num=_num) _entries = _conventionaize(wtfs[:, :_num]) return _constructors.frequency_series(freqs, entries=_entries)