commkit 1.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (84) hide show
  1. commkit/__init__.py +74 -0
  2. commkit/_cuda/__init__.py +321 -0
  3. commkit/_cuda/compiler.py +88 -0
  4. commkit/_cuda/src/bps_min_d2.cu +104 -0
  5. commkit/_cuda/src/cs_block.cu +119 -0
  6. commkit/_cuda/src/selftest.cu +14 -0
  7. commkit/analysis/__init__.py +55 -0
  8. commkit/analysis/_common.py +236 -0
  9. commkit/analysis/allan.py +108 -0
  10. commkit/analysis/drift.py +213 -0
  11. commkit/analysis/interferometry.py +887 -0
  12. commkit/analysis/linewidth.py +480 -0
  13. commkit/analysis/trajectory.py +91 -0
  14. commkit/backend.py +507 -0
  15. commkit/coding/__init__.py +23 -0
  16. commkit/coding/base.py +17 -0
  17. commkit/coding/bch.py +6 -0
  18. commkit/coding/convolutional.py +7 -0
  19. commkit/coding/crc.py +7 -0
  20. commkit/coding/galois.py +8 -0
  21. commkit/coding/hamming.py +6 -0
  22. commkit/coding/interleaving.py +7 -0
  23. commkit/coding/ldpc.py +8 -0
  24. commkit/coding/polar.py +8 -0
  25. commkit/coding/ratematch.py +6 -0
  26. commkit/coding/reed_solomon.py +6 -0
  27. commkit/coding/turbo.py +8 -0
  28. commkit/core/__init__.py +32 -0
  29. commkit/core/frame.py +992 -0
  30. commkit/core/generation.py +581 -0
  31. commkit/core/signal.py +725 -0
  32. commkit/equalization/__init__.py +49 -0
  33. commkit/equalization/_block.py +1855 -0
  34. commkit/equalization/_common.py +606 -0
  35. commkit/equalization/_kernels_jax.py +1720 -0
  36. commkit/equalization/_kernels_numba.py +1704 -0
  37. commkit/equalization/blind.py +223 -0
  38. commkit/equalization/linear.py +365 -0
  39. commkit/equalization/polarization.py +790 -0
  40. commkit/equalization/result.py +191 -0
  41. commkit/equalization/sequential.py +2805 -0
  42. commkit/filtering.py +1120 -0
  43. commkit/frequency.py +1191 -0
  44. commkit/helpers.py +489 -0
  45. commkit/impairments/__init__.py +43 -0
  46. commkit/impairments/channel/__init__.py +20 -0
  47. commkit/impairments/channel/linear.py +310 -0
  48. commkit/impairments/channel/nonlinear.py +11 -0
  49. commkit/impairments/frontend.py +229 -0
  50. commkit/impairments/noise.py +105 -0
  51. commkit/impairments/source.py +219 -0
  52. commkit/io.py +308 -0
  53. commkit/logger.py +103 -0
  54. commkit/mapping/__init__.py +46 -0
  55. commkit/mapping/bits.py +240 -0
  56. commkit/mapping/constellation.py +153 -0
  57. commkit/mapping/gray.py +429 -0
  58. commkit/mapping/llr.py +253 -0
  59. commkit/mapping/shaping.py +218 -0
  60. commkit/metrics.py +949 -0
  61. commkit/multirate.py +476 -0
  62. commkit/plotting/__init__.py +78 -0
  63. commkit/plotting/analysis.py +627 -0
  64. commkit/plotting/constellation.py +483 -0
  65. commkit/plotting/equalizer.py +390 -0
  66. commkit/plotting/eye.py +388 -0
  67. commkit/plotting/spectral.py +575 -0
  68. commkit/plotting/sync.py +953 -0
  69. commkit/plotting/theme.py +203 -0
  70. commkit/plotting/waveform.py +200 -0
  71. commkit/py.typed +0 -0
  72. commkit/recovery/__init__.py +51 -0
  73. commkit/recovery/bps.py +337 -0
  74. commkit/recovery/corrections.py +751 -0
  75. commkit/recovery/pilots.py +803 -0
  76. commkit/recovery/pll.py +482 -0
  77. commkit/recovery/tikhonov.py +424 -0
  78. commkit/recovery/viterbi_viterbi.py +227 -0
  79. commkit/spectral.py +560 -0
  80. commkit/timing.py +841 -0
  81. commkit-1.0.0.dist-info/METADATA +145 -0
  82. commkit-1.0.0.dist-info/RECORD +84 -0
  83. commkit-1.0.0.dist-info/WHEEL +4 -0
  84. commkit-1.0.0.dist-info/licenses/LICENSE +21 -0
commkit/spectral.py ADDED
@@ -0,0 +1,560 @@
1
+ """
2
+ Spectral analysis and frequency-domain processing.
3
+
4
+ This module provides high-performance routines for spectral estimation and
5
+ manipulation, optimized for both CPU and GPU backends. It supports Welch's
6
+ Power Spectral Density (PSD) method and phase-continuous frequency shifting.
7
+ """
8
+
9
+ from collections.abc import Sequence
10
+ from typing import Any, cast
11
+
12
+ import numpy as np
13
+
14
+ from .backend import ArrayType, dispatch
15
+ from .core.signal import Signal
16
+ from .logger import logger
17
+
18
+
19
+ def shift_frequency(
20
+ samples: ArrayType | Signal,
21
+ offset: float,
22
+ sampling_rate: float | None = None,
23
+ ) -> tuple[ArrayType, float] | Signal:
24
+ """
25
+ Applies a frequency offset (complex mixing) to a signal.
26
+
27
+ This function shifts the signal spectrum by a specified offset in Hz
28
+ by multiplying the samples with a complex phasor:
29
+ s_shifted(t) = s(t) * e^(j * 2 * pi * f_offset * t)
30
+
31
+ To maintain phase continuity and prevent spectral leakage when the
32
+ signal is treated as periodic (e.g., in circular convolution or
33
+ FFT-based operations), the applied offset is quantized to the
34
+ fundamental frequency resolution of the signal (df = f_s / N).
35
+
36
+ Parameters
37
+ ----------
38
+ samples : array_like or Signal
39
+ Input signal samples. Shape: (..., N_samples).
40
+ offset : float
41
+ Target frequency shift in Hz. Positive values shift the spectrum
42
+ towards higher frequencies.
43
+ sampling_rate : float
44
+ Sampling rate in Hz.
45
+
46
+ Returns
47
+ -------
48
+ shifted_samples : array_like
49
+ The frequency-shifted signal on the same backend as the input.
50
+ actual_offset : float
51
+ The actual quantized frequency shift applied to the signal.
52
+
53
+ Notes
54
+ -----
55
+ The quantization ensures that the applied shift corresponds to an
56
+ integer number of cycles over the signal duration, which is critical
57
+ for preserving the circularity of the signal's phase.
58
+
59
+ When ``samples`` is a :class:`Signal`, a new :class:`Signal` is returned
60
+ with the shift applied and ``digital_frequency_offset`` accumulated;
61
+ ``sampling_rate`` is taken from the signal.
62
+ """
63
+ if isinstance(samples, Signal):
64
+ sig = samples
65
+ new = sig.copy()
66
+ out = shift_frequency(sig.samples, offset, sig.sampling_rate)
67
+ assert isinstance(out, tuple) # array input -> (samples, actual_offset)
68
+ shifted, actual = out
69
+ new.samples = shifted
70
+ if new.digital_frequency_offset is None:
71
+ new.digital_frequency_offset = 0.0
72
+ new.digital_frequency_offset += actual
73
+ return new
74
+
75
+ if sampling_rate is None:
76
+ raise ValueError("shift_frequency() requires sampling_rate for array input.")
77
+
78
+ samples, xp, _ = dispatch(samples)
79
+
80
+ # Axis -1 is time
81
+ n = samples.shape[-1]
82
+ df = sampling_rate / n
83
+
84
+ # Quantize offset to nearest bin to ensure phase continuity
85
+ k = xp.round(offset / df)
86
+ actual_offset = k * df
87
+
88
+ if not xp.isclose(offset, actual_offset):
89
+ logger.warning(
90
+ "Requested offset %.3f Hz quantized to %.3f Hz (step %.3f Hz) to maintain phase continuity.",
91
+ offset,
92
+ actual_offset,
93
+ df,
94
+ )
95
+ else:
96
+ logger.debug("Applying frequency offset: %.3f Hz.", actual_offset)
97
+
98
+ # Time vector
99
+ t = xp.arange(n) / sampling_rate
100
+
101
+ # Apply mixing
102
+ # exp(j * 2 * pi * f * t)
103
+ # Phase is computed at float64 accuracy (xp.pi is float64), then the mixer
104
+ # is cast to the signal's complex precision to prevent silent promotion of
105
+ # complex64/float32 signals to complex128/float64.
106
+ phase = 2 * xp.pi * actual_offset * t
107
+ mixer = xp.exp(1j * phase) # complex128
108
+ if xp.iscomplexobj(samples):
109
+ target_cdtype = samples.dtype
110
+ else:
111
+ target_cdtype = xp.complex64 if samples.dtype == xp.float32 else xp.complex128
112
+ mixer = mixer.astype(target_cdtype)
113
+
114
+ # Broadcast mixer to match samples shape: (1, ..., 1, N)
115
+ if samples.ndim > 1:
116
+ mixer = mixer.reshape((1,) * (samples.ndim - 1) + (-1,))
117
+
118
+ return samples * mixer, float(actual_offset)
119
+
120
+
121
+ def add_pilot_tone(
122
+ samples: ArrayType | Signal,
123
+ sampling_rate: float | None = None,
124
+ frequency: float | Sequence[float] | None = None,
125
+ power_ratio_db: float | Sequence[float] = -15.0,
126
+ phase_init: float = 0.0,
127
+ renormalize: bool = False,
128
+ ) -> tuple[ArrayType, float | list[float]] | Signal:
129
+ r"""
130
+ Add a continuous-wave (CW) pilot tone to a baseband waveform.
131
+
132
+ Superimposes a * exp(j*(2*pi*f_p*n/f_s + phi_0)) on the oversampled samples.
133
+ The tone acquires the same carrier frequency offset and phase noise as the
134
+ data; at the receiver its phase directly recovers both - see
135
+ ``recover_carrier_phase_pilot_tone``.
136
+
137
+ Apply to a pulse-shaped oversampled waveform before channel impairments.
138
+ Place the tone in a guard band: (1+beta)/2 * R_s < |f_p| < f_s/2.
139
+
140
+ Parameters
141
+ ----------
142
+ samples : array_like
143
+ Complex baseband samples. Shape: ``(N,)`` (SISO) or ``(C, N)`` (MIMO).
144
+ The tone is added to every channel, each scaled to its own power.
145
+ frequency : float or sequence of float
146
+ Requested tone frequency fp in Hz, in ``(-f_s/2, f_s/2)``.
147
+ A **scalar** places the same tone on every channel. A **sequence**
148
+ of length ``C`` places one tone per channel (channel ``c`` gets
149
+ ``frequency[c]``) - distinct per-channel tones enable, e.g.,
150
+ tone-based polarization demultiplexing
151
+ (``demultiplex_polarization_tones``). Each value is quantized to the
152
+ nearest FFT bin ``f_s/N`` (see Notes); the **actual** applied
153
+ frequency(ies) are returned.
154
+ sampling_rate : float
155
+ Sampling rate fs in Hz.
156
+ power_ratio_db : float or sequence of float, default -15.0
157
+ Pilot-to-signal power ratio (PSR) in dB: 10*log10(P_tone / P_signal).
158
+ Typical range -20 to -10 dB. A **scalar** applies the same PSR to every
159
+ channel; a **sequence** of length ``C`` sets one PSR per channel
160
+ (mirroring per-channel ``frequency``).
161
+ phase_init : float, default 0.0
162
+ Initial tone phase phi_0 in radians, common to all channels. Acts as
163
+ a known phase reference; it appears as a constant offset in the
164
+ recovered phase and is absorbed by the usual post-CPR ambiguity
165
+ resolution.
166
+ renormalize : bool, default False
167
+ If ``True``, rescale each channel after adding the tone so its mean
168
+ power matches the input (preserves the library power invariant E[|x|²] = 1/sps).
169
+ If ``False``, total power rises by 1 + 10^(PSR/10).
170
+
171
+ Returns
172
+ -------
173
+ samples : array_like
174
+ Samples with the pilot tone added, same shape, dtype, and backend as
175
+ the input.
176
+ actual_frequency : float or list of float
177
+ The grid-quantized tone frequency(ies) in Hz actually applied (see
178
+ Notes). A **scalar** ``frequency`` returns a single ``float``; a
179
+ per-channel **sequence** returns a ``list`` of ``C`` floats. Store
180
+ this (e.g. in ``pilot_tone_frequency``) and pass it to the receiver,
181
+ since it - not the requested value - is where the tone(s) sit.
182
+
183
+ Raises
184
+ ------
185
+ ValueError
186
+ If any requested frequency lies outside ``(-fs/2, fs/2)``, or if a
187
+ per-channel sequence is given whose length does not equal ``C``.
188
+
189
+ Notes
190
+ -----
191
+ Each requested frequency is snapped to the nearest FFT bin (fs/N) so the
192
+ tone completes an integer number of cycles per buffer - ensuring seamless
193
+ playback on an AWG/DAC. The quantization error is at most fs/(2N).
194
+ The phase ramp is accumulated in float64 to avoid trig argument-reduction
195
+ error for large N.
196
+
197
+ When ``samples`` is a :class:`Signal`, the sampling rate is taken from the
198
+ signal, so the **second positional argument is the frequency** (i.e. call
199
+ ``add_pilot_tone(sig, freq, ...)``). A new :class:`Signal` is returned with
200
+ ``pilot_tone_frequency`` / ``pilot_tone_power_ratio_db`` recorded.
201
+ """
202
+ if isinstance(samples, Signal):
203
+ sig = samples
204
+ # Signal rate is implicit; the second positional carries the frequency.
205
+ freq = frequency if frequency is not None else sampling_rate
206
+ if freq is None:
207
+ raise ValueError("add_pilot_tone() requires a frequency.")
208
+ new = sig.copy()
209
+ out = add_pilot_tone(
210
+ sig.samples,
211
+ sig.sampling_rate,
212
+ freq,
213
+ power_ratio_db=power_ratio_db,
214
+ phase_init=phase_init,
215
+ renormalize=renormalize,
216
+ )
217
+ assert isinstance(out, tuple) # array input -> (samples, actual_frequency)
218
+ new.samples, new.pilot_tone_frequency = out
219
+ new.pilot_tone_power_ratio_db = power_ratio_db
220
+ return new
221
+
222
+ if sampling_rate is None or frequency is None:
223
+ raise ValueError(
224
+ "add_pilot_tone() requires sampling_rate and frequency for array input."
225
+ )
226
+
227
+ samples, xp, _ = dispatch(samples)
228
+ was_1d = samples.ndim == 1
229
+ if was_1d:
230
+ samples = samples[None, :] # (1, N)
231
+ C, N = samples.shape
232
+
233
+ # Normalise ``frequency`` to a per-channel (C,) host array. A scalar is
234
+ # broadcast to every channel (and returns a scalar for back-compat); a
235
+ # sequence must supply exactly one frequency per channel.
236
+ scalar_input = np.ndim(frequency) == 0
237
+ if scalar_input:
238
+ f_req = [float(cast(float, frequency))] * C
239
+ else:
240
+ f_req = [float(f) for f in cast(Sequence[float], frequency)]
241
+ if len(f_req) != C:
242
+ raise ValueError(
243
+ f"frequency sequence has length {len(f_req)} but the signal has "
244
+ f"C={C} channel(s); supply one frequency per channel."
245
+ )
246
+
247
+ nyq = sampling_rate / 2.0
248
+ for f in f_req:
249
+ if not (-nyq < f < nyq):
250
+ raise ValueError(
251
+ f"frequency={f} must lie in (-fs/2, fs/2) = (±{nyq:.3g}) Hz."
252
+ )
253
+
254
+ # Snap each tone to the FFT bin grid so it is buffer-periodic (loop-seamless
255
+ # on an AWG/DAC), mirroring shift_frequency's quantization.
256
+ df = sampling_rate / N
257
+ actual = [float(round(f / df) * df) for f in f_req]
258
+ for f_in, f_out in zip(f_req, actual):
259
+ if abs(f_out - f_in) > 1e-12 * max(1.0, abs(f_in)):
260
+ logger.warning(
261
+ "add_pilot_tone: requested %.3f Hz quantized to %.3f Hz (grid step fs/N=%.3f Hz) for buffer-periodic (loop-seamless) playback.",
262
+ f_in,
263
+ f_out,
264
+ df,
265
+ )
266
+
267
+ # Normalise ``power_ratio_db`` to a per-channel (C,) list, mirroring how
268
+ # ``frequency`` is handled: a scalar broadcasts to every channel; a sequence
269
+ # must supply exactly one PSR per channel.
270
+ scalar_power = np.ndim(power_ratio_db) == 0
271
+ if scalar_power:
272
+ psr_req = [float(cast(float, power_ratio_db))] * C
273
+ else:
274
+ psr_req = [float(p) for p in cast(Sequence[float], power_ratio_db)]
275
+ if len(psr_req) != C:
276
+ raise ValueError(
277
+ f"power_ratio_db sequence has length {len(psr_req)} but the signal "
278
+ f"has C={C} channel(s); supply one PSR per channel."
279
+ )
280
+
281
+ # Per-channel signal power and the tone amplitude that realises the PSR.
282
+ p_signal = xp.mean(xp.abs(samples) ** 2, axis=-1, keepdims=True) # (C, 1) float
283
+ psr_lin = (10.0 ** (xp.asarray(psr_req, dtype=xp.float64) / 10.0)).reshape(
284
+ C, 1
285
+ ) # (C, 1)
286
+ amp = xp.sqrt(p_signal * psr_lin) # (C, 1)
287
+
288
+ # Per-channel phase ramp (C, N) in float64; wrap to [-π, π) before exp so
289
+ # complex64 targets avoid argument-reduction error on long ramps
290
+ # (cf. correct_static_frequency_offset).
291
+ two_pi = 2.0 * xp.pi
292
+ n = xp.arange(N, dtype=xp.float64) # (N,)
293
+ f_ch = xp.asarray(actual, dtype=xp.float64).reshape(C, 1) # (C, 1)
294
+ phase = two_pi * f_ch * n[None, :] / sampling_rate + phase_init # (C, N) float64
295
+ phase = phase - xp.round(phase / two_pi) * two_pi
296
+
297
+ dtype_real = xp.float32 if samples.dtype == xp.complex64 else xp.float64
298
+ tone = xp.exp(1j * phase.astype(dtype_real)).astype(samples.dtype) # (C, N)
299
+ out = samples + amp.astype(samples.dtype) * tone # (C, N)
300
+
301
+ if renormalize:
302
+ # Restore each channel to its original mean power.
303
+ p_out = xp.mean(xp.abs(out) ** 2, axis=-1, keepdims=True) # (C, 1)
304
+ out = out * xp.sqrt(p_signal / p_out).astype(samples.dtype)
305
+
306
+ f_log = f"{actual[0]:.3g} Hz" if scalar_input else f"{actual} Hz"
307
+ psr_log = f"{psr_req[0]:.1f} dB" if scalar_power else f"{psr_req} dB"
308
+ logger.info(
309
+ "add_pilot_tone: f_p=%s, PSR=%s, phase_init=%.3g rad, renormalize=%s [C=%s, N=%s]",
310
+ f_log,
311
+ psr_log,
312
+ phase_init,
313
+ renormalize,
314
+ C,
315
+ N,
316
+ )
317
+
318
+ out = out[0] if was_1d else out
319
+ actual_frequency: float | list[float] = actual[0] if scalar_input else actual
320
+ return out, actual_frequency
321
+
322
+
323
+ def welch_psd(
324
+ samples: ArrayType | Signal,
325
+ sampling_rate: float | None = None,
326
+ nperseg: int = 256,
327
+ detrend: str | bool | None = False,
328
+ average: str | None = "mean",
329
+ window: str | tuple[Any, ...] | Any = "hann",
330
+ noverlap: int | None = None,
331
+ nfft: int | None = None,
332
+ scaling: str = "density",
333
+ return_onesided: bool | None = None,
334
+ axis: int = -1,
335
+ ) -> tuple[ArrayType, ArrayType]:
336
+ """
337
+ Estimates the Power Spectral Density (PSD) using Welch's method.
338
+
339
+ Welch's method provides a lower-variance estimate of the PSD
340
+ compared to a raw periodogram by averaging spectra computed over
341
+ overlapping segments of the signal.
342
+
343
+ Parameters
344
+ ----------
345
+ samples : array_like or Signal
346
+ Input signal samples. Shape: (..., N_samples).
347
+ sampling_rate : float
348
+ Sampling rate in Hz.
349
+ nperseg : int, default 256
350
+ Length of each segment. A longer segment increases frequency
351
+ resolution but also increases the variance of the estimate.
352
+ detrend : str or bool, default False
353
+ Specifies how to detrend each segment (e.g., 'constant', 'linear').
354
+ average : {"mean", "median"}, default "mean"
355
+ Method to use for averaging segments. Median is more robust to
356
+ transient outliers.
357
+ window : str or tuple or array_like, default "hann"
358
+ Desired window to use. If `window` is a string or tuple, it is
359
+ passed to `scipy.signal.get_window` to generate the window values.
360
+ noverlap : int, optional
361
+ Number of points to overlap between segments. If None,
362
+ `noverlap = nperseg // 2`.
363
+ nfft : int, optional
364
+ Length of the FFT used, if a zero padded FFT is desired. If None,
365
+ the FFT length is `nperseg`.
366
+ scaling : {"density", "spectrum"}, default "density"
367
+ Selects between computing the power spectral density ('density')
368
+ where Pxx has units of V**2/Hz and computing the power spectrum
369
+ ('spectrum') where Pxx has units of V**2.
370
+ return_onesided : bool, optional
371
+ If True, returns a one-sided spectrum (frequencies 0 to f_s/2)
372
+ for real-valued data. For complex data, only two-sided spectra
373
+ (frequencies -f_s/2 to f_s/2) are supported.
374
+ Axis along which to compute the PSD.
375
+
376
+ Returns
377
+ -------
378
+ f : array_like
379
+ Array of sample frequencies.
380
+ Pxx : array_like
381
+ Power spectral density (linear scale, units: V^2/Hz).
382
+
383
+ Raises
384
+ ------
385
+ ValueError
386
+ If `return_onesided` set to True for complex-valued inputs.
387
+ """
388
+ if isinstance(samples, Signal):
389
+ sig = samples
390
+ return welch_psd(
391
+ sig.samples,
392
+ sig.sampling_rate,
393
+ nperseg=nperseg,
394
+ detrend=detrend,
395
+ average=average,
396
+ window=window,
397
+ noverlap=noverlap,
398
+ nfft=nfft,
399
+ scaling=scaling,
400
+ return_onesided=return_onesided,
401
+ axis=-1,
402
+ )
403
+
404
+ if sampling_rate is None:
405
+ raise ValueError("welch_psd() requires sampling_rate for array input.")
406
+
407
+ samples, xp, sp = dispatch(samples)
408
+ is_complex = xp.iscomplexobj(samples)
409
+
410
+ if return_onesided is None:
411
+ return_onesided = not is_complex
412
+
413
+ # scipy.signal.welch returns onesided by default for real, two-sided for complex
414
+ # unless return_onesided is explicitly set.
415
+ # Note: scipy's return_onesided argument serves to force one-sided for real data.
416
+ # It cannot force one-sided for complex data (always raises error).
417
+ # For complex data, it always returns two-sided (0 to fs).
418
+
419
+ if is_complex and return_onesided:
420
+ raise ValueError("Cannot compute one-sided PSD for complex data.")
421
+
422
+ f, Pxx = sp.signal.welch(
423
+ samples,
424
+ fs=sampling_rate,
425
+ window=window,
426
+ nperseg=nperseg,
427
+ noverlap=noverlap,
428
+ nfft=nfft,
429
+ detrend=detrend,
430
+ return_onesided=return_onesided,
431
+ scaling=scaling,
432
+ axis=axis,
433
+ average=average,
434
+ )
435
+
436
+ if not return_onesided:
437
+ # Shift zero frequency to center
438
+ # f is typically 1D array of frequencies
439
+ f = xp.fft.fftshift(f)
440
+ # Pxx needs shift along the frequency axis
441
+ Pxx = xp.fft.fftshift(Pxx, axes=axis)
442
+
443
+ return f, Pxx
444
+
445
+
446
+ def spectrogram(
447
+ samples: ArrayType | Signal,
448
+ sampling_rate: float | None = None,
449
+ window: str | tuple[Any, ...] | Any = "hann",
450
+ nperseg: int = 256,
451
+ noverlap: int | None = None,
452
+ nfft: int | None = None,
453
+ detrend: str | bool | None = False,
454
+ return_onesided: bool | None = None,
455
+ scaling: str = "density",
456
+ axis: int = -1,
457
+ mode: str = "psd",
458
+ ) -> tuple[ArrayType, ArrayType, ArrayType]:
459
+ """
460
+ Computes a spectrogram with consecutive Fourier transforms.
461
+
462
+ Parameters
463
+ ----------
464
+ samples : array_like or Signal
465
+ Input signal samples. Shape: (..., N_samples).
466
+ sampling_rate : float
467
+ Sampling rate in Hz.
468
+ window : str or tuple or array_like, default "hann"
469
+ Desired window to use. If `window` is a string or tuple, it is
470
+ passed to `scipy.signal.get_window` to generate the window values.
471
+ nperseg : int, default 256
472
+ Length of each segment. A longer segment increases frequency
473
+ resolution but also increases the variance of the estimate.
474
+ noverlap : int, optional
475
+ Number of points to overlap between segments. If None,
476
+ `noverlap = nperseg // 2`.
477
+ nfft : int, optional
478
+ Length of the FFT used, if a zero padded FFT is desired. If None,
479
+ the FFT length is `nperseg`.
480
+ detrend : str or bool, default False
481
+ Specifies how to detrend each segment (e.g., 'constant', 'linear').
482
+ return_onesided : bool, optional
483
+ If True, returns a one-sided spectrum (frequencies 0 to f_s/2)
484
+ for real-valued data. For complex data, only two-sided spectra
485
+ are supported.
486
+ scaling : {"density", "spectrum"}, default "density"
487
+ Selects between computing the power spectral density ('density')
488
+ where Sxx has units of V**2/Hz and computing the power spectrum
489
+ ('spectrum') where Sxx has units of V**2.
490
+ axis : int, default -1
491
+ The axis along which to compute the spectrogram.
492
+ mode : {"psd", "complex", "magnitude", "angle", "phase"}, default "psd"
493
+ Type of spectrogram to return. Options are 'psd', 'complex',
494
+ 'magnitude', 'angle', 'phase'.
495
+
496
+ Returns
497
+ -------
498
+ f : array_like
499
+ Array of sample frequencies.
500
+ t : array_like
501
+ Array of segment times.
502
+ Sxx : array_like
503
+ Spectrogram of the signal.
504
+
505
+ Raises
506
+ ------
507
+ ValueError
508
+ If `return_onesided` set to True for complex-valued inputs.
509
+ """
510
+ if isinstance(samples, Signal):
511
+ sig = samples
512
+ return spectrogram(
513
+ sig.samples,
514
+ sig.sampling_rate,
515
+ window=window,
516
+ nperseg=nperseg,
517
+ noverlap=noverlap,
518
+ nfft=nfft,
519
+ detrend=detrend,
520
+ return_onesided=return_onesided,
521
+ scaling=scaling,
522
+ axis=-1,
523
+ mode=mode,
524
+ )
525
+
526
+ if sampling_rate is None:
527
+ raise ValueError("spectrogram() requires sampling_rate for array input.")
528
+
529
+ samples, xp, sp = dispatch(samples)
530
+ is_complex = xp.iscomplexobj(samples)
531
+
532
+ if return_onesided is None:
533
+ return_onesided = not is_complex
534
+
535
+ if is_complex and return_onesided:
536
+ raise ValueError("Cannot compute one-sided spectrogram for complex data.")
537
+
538
+ f, t, Sxx = sp.signal.spectrogram(
539
+ samples,
540
+ fs=sampling_rate,
541
+ window=window,
542
+ nperseg=nperseg,
543
+ noverlap=noverlap,
544
+ nfft=nfft,
545
+ detrend=detrend,
546
+ return_onesided=return_onesided,
547
+ scaling=scaling,
548
+ axis=axis,
549
+ mode=mode,
550
+ )
551
+
552
+ if not return_onesided:
553
+ # Shift zero frequency to center
554
+ f = xp.fft.fftshift(f)
555
+ # Sxx frequency axis is at position axis_pos in output
556
+ ndim = samples.ndim
557
+ axis_pos = axis % ndim
558
+ Sxx = xp.fft.fftshift(Sxx, axes=axis_pos)
559
+
560
+ return f, t, Sxx