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
@@ -0,0 +1,803 @@
1
+ """Pilot-symbol and pilot-tone aided carrier phase recovery."""
2
+
3
+ import logging
4
+
5
+ import numpy as np
6
+
7
+ from ..backend import ArrayType, dispatch, to_device
8
+ from ..logger import logger
9
+ from .corrections import correct_cycle_slips
10
+
11
+
12
+ def recover_carrier_phase_pilot_symbols(
13
+ symbols: ArrayType,
14
+ pilot_indices: ArrayType,
15
+ pilot_values: ArrayType,
16
+ interpolation: str = "linear",
17
+ joint_channels: bool = False,
18
+ cycle_slip_correction: bool = False,
19
+ cycle_slip_history: int = 100,
20
+ cycle_slip_threshold: float = np.pi / 4,
21
+ debug_plot: bool = False,
22
+ ) -> ArrayType:
23
+ """
24
+ Carrier phase recovery using known pilot symbols.
25
+
26
+ Computes the phase error at each pilot position, unwraps the pilot
27
+ phase sequence, and interpolates across the full symbol grid.
28
+
29
+ Parameters
30
+ ----------
31
+ symbols : array_like
32
+ Received 1-SPS complex symbols. Shape: (N,) or (C, N).
33
+ pilot_indices : array_like of int
34
+ Indices of pilot symbols within the frame, in increasing order.
35
+ Shape: (P,).
36
+ pilot_values : array_like
37
+ Known transmitted pilot constellation points.
38
+ Shape: (P,) for shared pilots (broadcast to all MIMO channels),
39
+ or (C, P) for per-channel pilots.
40
+ interpolation : {'linear', 'cubic'}, default 'linear'
41
+ Interpolation method between pilot positions. Both modes loop over
42
+ MIMO channels (``xp.interp`` and ``CubicSpline`` are 1D-only);
43
+ C is typically 1-4 so the overhead is negligible. ``'cubic'`` uses
44
+ ``CubicSpline`` (CPU) or
45
+ ``CubicSpline`` (GPU) with natural
46
+ boundary conditions (zero second derivative at endpoints) and
47
+ constant-hold extrapolation outside the pilot span.
48
+ joint_channels : bool, default False
49
+ For MIMO inputs (C > 1): if ``True``, perform coherent complex
50
+ averaging of ``r_pilot * conj(s_pilot)`` across all channels before
51
+ calling ``angle()``. This avoids wrap-around artefacts that arise
52
+ when averaging phases directly, and reduces variance by ~√C for
53
+ shared-LO systems. The resulting single phase trajectory is broadcast
54
+ to all C output rows. Has no effect for SISO (C = 1).
55
+ cycle_slip_correction : bool, default False
56
+ If ``True``, apply ``correct_cycle_slips`` to the unwrapped pilot
57
+ phase sequence before interpolation, with ``symmetry=1`` (correction
58
+ quantum ``2π``) to detect and fix wrap-around errors introduced by
59
+ ``xp.unwrap`` at large inter-pilot gaps.
60
+ cycle_slip_history : int, default 100
61
+ ``history_length`` passed to ``correct_cycle_slips``.
62
+ cycle_slip_threshold : float, default π/4
63
+ ``threshold`` passed to ``correct_cycle_slips`` (radians).
64
+ debug_plot : bool, default False
65
+ If ``True``, opens a diagnostic figure showing the unwrapped pilot
66
+ phase sequence and the interpolated phase trajectory.
67
+
68
+ Returns
69
+ -------
70
+ array_like
71
+ Per-symbol phase estimate in radians. Shape matches ``symbols``.
72
+ Same backend as input.
73
+
74
+ Notes
75
+ -----
76
+ Phase at each pilot: phi_hat[k] = angle(r[k] * conj(s[k])). Linear
77
+ interpolation constant-holds at the boundaries; cubic uses natural spline
78
+ with constant-hold extrapolation. Single-carrier only.
79
+ """
80
+ symbols, xp, _ = dispatch(symbols)
81
+ was_1d = symbols.ndim == 1
82
+ if was_1d:
83
+ symbols = symbols[None, :]
84
+ C, N = symbols.shape
85
+
86
+ pilot_indices_np = to_device(pilot_indices, "cpu").astype(np.intp)
87
+ pilot_indices_xp = xp.asarray(pilot_indices, dtype=xp.float64)
88
+ pilot_values_xp = xp.asarray(pilot_values)
89
+ P = len(pilot_indices_np)
90
+
91
+ # Broadcast shared pilots (P,) -> (C, P) for all channels
92
+ if pilot_values_xp.ndim == 1:
93
+ pilot_values_xp = xp.broadcast_to(pilot_values_xp[None, :], (C, P))
94
+
95
+ # Phase at each pilot position: angle(r_pilot · conj(s_pilot))
96
+ r_pilots = symbols[:, pilot_indices_np] # (C, P)
97
+
98
+ if joint_channels and C > 1:
99
+ # Coherent complex averaging before angle() - avoids wrap-around artefacts
100
+ # that arise from averaging phases directly (e.g. antipodal channels).
101
+ z_joint = xp.mean(r_pilots * xp.conj(pilot_values_xp), axis=0) # (P,)
102
+ phi_joint_u = xp.unwrap(xp.angle(z_joint).astype(xp.float64)) # (P,)
103
+ if cycle_slip_correction:
104
+ phi_joint_np = to_device(phi_joint_u, "cpu")
105
+ phi_joint_np = correct_cycle_slips(
106
+ phi_joint_np,
107
+ symmetry=1,
108
+ history_length=cycle_slip_history,
109
+ threshold=cycle_slip_threshold,
110
+ )
111
+ phi_joint_u = xp.asarray(phi_joint_np)
112
+ # Broadcast to (C, P) - read-only, downstream code only reads phi_pilots_u[ch]
113
+ phi_pilots_u = xp.broadcast_to(phi_joint_u[None, :], (C, P))
114
+ else:
115
+ phi_pilots = xp.angle(r_pilots * xp.conj(pilot_values_xp)) # (C, P)
116
+ # Unwrap along the pilot axis in float64 (cp.unwrap preserves input dtype;
117
+ # casting before avoids precision loss in the discontinuity test for float32 input)
118
+ phi_pilots_u = xp.unwrap(phi_pilots.astype(xp.float64), axis=-1) # (C, P)
119
+ if cycle_slip_correction:
120
+ phi_pilots_u_np = to_device(phi_pilots_u, "cpu")
121
+ for ch in range(C):
122
+ phi_pilots_u_np[ch] = correct_cycle_slips(
123
+ phi_pilots_u_np[ch],
124
+ symmetry=1,
125
+ history_length=cycle_slip_history,
126
+ threshold=cycle_slip_threshold,
127
+ )
128
+ phi_pilots_u = xp.asarray(phi_pilots_u_np)
129
+
130
+ all_positions = xp.arange(N, dtype=xp.float64)
131
+
132
+ if interpolation == "linear":
133
+ # xp.interp handles non-uniform pilot spacing natively, is boundary-safe
134
+ # (extrapolates with first/last pilot value), and avoids the divide-by-zero
135
+ # guards that the searchsorted form required. Loop over C channels because
136
+ # xp.interp is 1D-only; overhead is negligible for typical C = 1-4.
137
+ phi_full = xp.empty((C, N), dtype=xp.float64)
138
+ for ch in range(C):
139
+ phi_full[ch] = xp.interp(all_positions, pilot_indices_xp, phi_pilots_u[ch])
140
+
141
+ elif interpolation == "cubic":
142
+ # CubicSpline is inherently per-channel (1D y input); loop is unavoidable.
143
+ # Both scipy (CPU) and cupyx.scipy (GPU) share the same API.
144
+ phi_full = xp.empty((C, N), dtype=xp.float64)
145
+ if xp is not np:
146
+ from cupyx.scipy.interpolate import CubicSpline
147
+ else:
148
+ from scipy.interpolate import CubicSpline
149
+
150
+ for ch in range(C):
151
+ phi_ch = phi_pilots_u[ch] # already float64
152
+ cs = CubicSpline(pilot_indices_xp, phi_ch, bc_type="natural")
153
+ # Evaluate the spline only within the pilot span; constant-hold outside.
154
+ first_idx = int(pilot_indices_np[0])
155
+ last_idx = int(pilot_indices_np[-1])
156
+ phi_full[ch, first_idx : last_idx + 1] = cs(
157
+ all_positions[first_idx : last_idx + 1]
158
+ )
159
+ if first_idx > 0:
160
+ phi_full[ch, :first_idx] = phi_ch[0]
161
+ if last_idx < N - 1:
162
+ phi_full[ch, last_idx + 1 :] = phi_ch[-1]
163
+
164
+ else:
165
+ raise ValueError(
166
+ f"Unknown interpolation method: {interpolation!r}. "
167
+ "Choose 'linear' or 'cubic'."
168
+ )
169
+
170
+ # Host copy of the trajectory is needed only for the INFO summary and the
171
+ # optional debug plot; skip the transfer + reductions otherwise (the device
172
+ # phi_full drives the actual correction and is what gets returned).
173
+ _want_log = logger.isEnabledFor(logging.INFO)
174
+ if _want_log or debug_plot:
175
+ phi_full_np = to_device(phi_full, "cpu")
176
+ if _want_log:
177
+ phi_mean_deg = float(np.mean(phi_full_np)) * 180.0 / np.pi
178
+ phi_std_deg = float(np.std(phi_full_np)) * 180.0 / np.pi
179
+ logger.info(
180
+ "CPR (pilot-aided, %s): phase mean=%.2f°, std=%.2f° [P=%s pilots, C=%s]",
181
+ interpolation,
182
+ phi_mean_deg,
183
+ phi_std_deg,
184
+ P,
185
+ C,
186
+ )
187
+
188
+ if debug_plot:
189
+ from .. import plotting as _plotting
190
+
191
+ phi_pilots_u_np = to_device(phi_pilots_u, "cpu")
192
+ _plotting.plot_pilot_phase_estimate(
193
+ pilot_indices=pilot_indices_np,
194
+ phi_pilots_u=phi_pilots_u_np,
195
+ phi_full=phi_full_np,
196
+ show=True,
197
+ title="CPR - Pilot-Aided Phase",
198
+ )
199
+
200
+ if was_1d:
201
+ return phi_full[0]
202
+ return phi_full
203
+
204
+
205
+ def _extract_pilot_phasor(
206
+ samples: ArrayType,
207
+ sampling_rate: float,
208
+ tone_frequency: float,
209
+ bandwidth: float,
210
+ xp,
211
+ search_band: float | None = None,
212
+ refine_tone: bool = True,
213
+ window: str | tuple = "tukey",
214
+ X: ArrayType | None = None,
215
+ return_window: bool = False,
216
+ ) -> tuple[ArrayType, np.ndarray, np.ndarray, np.ndarray, ArrayType | None, ArrayType]:
217
+ """Isolate a CW pilot tone and return its carrier-stripped complex phasor.
218
+
219
+ Shared core of the pilot-tone CPR functions
220
+ (``recover_carrier_phase_pilot_tone``, ``recover_carrier_phase_pilot_tones``):
221
+ refine the per-channel tone centre, extract it with a zero-phase spectral
222
+ window (FFT -> window -> IFFT, sample-aligned), and strip the nominal carrier
223
+ so a residual frequency offset survives as a slow phase ramp.
224
+
225
+ Everything runs in the signal's working precision (complex64 for complex64
226
+ input) - the ±π-safe part of the pipeline is the float64 promotion of the
227
+ *angle* before unwrap, which the callers already perform (CLAUDE.md), not
228
+ double-precision spectra. Tone refinement reuses the extraction FFT
229
+ (device-side log-parabolic fit; one host transfer) instead of running a
230
+ zero-padded full-record FFT per channel, and the window/noise statistics
231
+ touch only the ``O(bandwidth/df)`` occupied bins instead of dense (C, N)
232
+ arrays.
233
+
234
+ Parameters
235
+ ----------
236
+ samples : (C, N) array
237
+ Oversampled complex samples, already 2-D (caller handles the 1-D case).
238
+ xp : module
239
+ The dispatched array module for ``samples`` (numpy/cupy).
240
+ X : (C, N) array, optional
241
+ Precomputed ``xp.fft.fft(samples, axis=-1)`` in working precision.
242
+ Pass it when extracting several tones from the same record so the
243
+ record is transformed once (``recover_carrier_phase_pilot_tones``).
244
+ return_window : bool, default False
245
+ Build and return the dense (C, N) extraction window ``W`` (diagnostics
246
+ only); when ``False`` the ``W`` slot in the return tuple is ``None``.
247
+ (others) : see ``recover_carrier_phase_pilot_tone``.
248
+
249
+ Returns
250
+ -------
251
+ phasor : (C, N) complex, working precision
252
+ ``z(n) ≈ A·e^{jθ(n)}`` per channel (carrier-frequency stripped).
253
+ f_centers : (C,) float64
254
+ Detected per-channel tone centre [Hz].
255
+ sig_power : (C,) float64
256
+ In-band tone power ``|A|²`` within the tracking window, in the same
257
+ units as ``mean(|z|²)`` (for SNR weights).
258
+ noise_power : (C,) float64
259
+ Additive-noise power ``σ²`` within the tracking window, same units.
260
+ W : (C, N) float64 or None
261
+ The extraction window (only if ``return_window=True``).
262
+ X : (C, N) complex, working precision
263
+ The full FFT (reusable for further tones / diagnostics).
264
+ """
265
+ from ..frequency import _refine_tones_from_spectrum, correct_static_frequency_offset
266
+
267
+ C, N = samples.shape
268
+ df = sampling_rate / N
269
+ if search_band is None:
270
+ search_band = bandwidth
271
+
272
+ # 1) One FFT in working precision (nfft = N keeps the IFFT sample-aligned).
273
+ if X is None:
274
+ xw = (
275
+ samples
276
+ if samples.dtype == xp.complex128
277
+ else samples.astype(xp.complex64, copy=False)
278
+ )
279
+ X = xp.fft.fft(xw, axis=-1) # (C, N)
280
+ real_dtype = xp.float64 if X.dtype == xp.complex128 else xp.float32
281
+
282
+ # 2) Per-channel tone centre. Refinement absorbs a frequency offset that
283
+ # has dragged the tone away from nominal, so the window stays centred on
284
+ # it. Runs on the shared spectrum - no extra FFTs, one host transfer.
285
+ if refine_tone:
286
+ f_centers = _refine_tones_from_spectrum(
287
+ X,
288
+ sampling_rate,
289
+ [float(tone_frequency)] * C,
290
+ search_band,
291
+ rows=range(C),
292
+ )
293
+ else:
294
+ f_centers = np.full(C, float(tone_frequency), dtype=np.float64)
295
+
296
+ # 3) Zero-phase extraction window placed circularly at each channel's tone bin.
297
+ from scipy.signal import get_window
298
+
299
+ half = int(bandwidth // df) # bins from centre to band edge
300
+ n_win = 2 * half + 1
301
+ try:
302
+ win_cpu = np.asarray(get_window(window, n_win, fftbins=False), dtype=np.float64)
303
+ except (ValueError, TypeError) as exc:
304
+ raise ValueError(
305
+ f"Invalid window {window!r}: {exc}. Pass any scipy.signal.get_window "
306
+ "spec, e.g. 'tukey', ('tukey', 0.3), 'boxcar', ('gaussian', 50)."
307
+ ) from exc
308
+ k_centers = np.round(f_centers / df).astype(np.int64) % N # (C,) centre bins
309
+ idx_np = (k_centers[:, None] + np.arange(-half, half + 1)[None, :]) % N
310
+ idx = xp.asarray(idx_np) # (C, n_win) circular in-band bins
311
+ rows = xp.arange(C)[:, None]
312
+ Xb = X[rows, idx] # (C, n_win) gathered in-band spectrum
313
+
314
+ # Per-channel in-band tone power and additive-noise power, via Parseval
315
+ # (mean|z|² = ΣΣ|X·W|²/N²). The noise floor is the median |X|² of a guard
316
+ # band one window-width outside the passband (local, so a neighbouring tone
317
+ # on the other channel does not inflate it). Batched over channels on the
318
+ # gathered bins only - no dense (C, N) window/power arrays, one transfer.
319
+ win64 = xp.asarray(win_cpu) # (n_win,) float64 on device
320
+ pow_b = xp.abs(Xb).astype(xp.float64) ** 2 # (C, n_win)
321
+ guard_off = np.arange(half + 1, half + 1 + n_win)
322
+ guard_np = (
323
+ np.concatenate(
324
+ [
325
+ (k_centers[:, None] + guard_off[None, :]),
326
+ (k_centers[:, None] - guard_off),
327
+ ],
328
+ axis=1,
329
+ )
330
+ % N
331
+ ) # (C, 2·n_win)
332
+ pow_guard = xp.abs(X[rows, xp.asarray(guard_np)]).astype(xp.float64) ** 2
333
+ floor_psd = xp.median(pow_guard, axis=-1) # (C,)
334
+ noise_dev = floor_psd * (n_win / (N * N))
335
+ win_tot = xp.sum(pow_b * win64**2, axis=-1) # (C,)
336
+ sig_dev = xp.maximum(win_tot / (N * N) - noise_dev, 1e-30)
337
+ stats = to_device(xp.stack([sig_dev, noise_dev]), "cpu") # one transfer
338
+ sig_power, noise_power = stats[0], stats[1]
339
+
340
+ # 4) Windowed band -> time domain: scatter the weighted bins into an
341
+ # otherwise-zero spectrum (equivalent to the dense X·W, without the
342
+ # full-record multiply) and IFFT.
343
+ Xw = xp.zeros((C, N), dtype=X.dtype)
344
+ Xw[rows, idx] = Xb * win64.astype(real_dtype)[None, :]
345
+ tone_t = xp.fft.ifft(Xw, axis=-1) # (C, N) working precision
346
+
347
+ # 5) Strip the *nominal* carrier so a residual frequency offset survives as
348
+ # a phase ramp - exact (non-quantized) complex mixing, same primitive
349
+ # the FOE correctors use (float64 phase ramp, wrapped, then cast to
350
+ # tone_t's working precision).
351
+ phasor = correct_static_frequency_offset(tone_t, sampling_rate, tone_frequency)
352
+
353
+ W = None
354
+ if return_window: # dense window for diagnostics/plots only
355
+ W = xp.zeros((C, N), dtype=xp.float64)
356
+ W[rows, idx] = win64[None, :]
357
+ return phasor, f_centers, sig_power, noise_power, W, X
358
+
359
+
360
+ def recover_carrier_phase_pilot_tone(
361
+ samples: ArrayType,
362
+ sampling_rate: float,
363
+ tone_frequency: float,
364
+ bandwidth: float,
365
+ search_band: float | None = None,
366
+ refine_tone: bool = True,
367
+ window: str | tuple = "tukey",
368
+ remove_frequency_offset: bool = True,
369
+ joint_channels: bool = False,
370
+ debug_plot: bool = False,
371
+ ) -> ArrayType:
372
+ r"""
373
+ Carrier phase recovery from a continuous-wave (CW) pilot tone.
374
+
375
+ Reads the common carrier phase straight off a pilot tone added at the
376
+ transmitter (see ``add_pilot_tone``). Because
377
+ the tone shares the data's local oscillator and channel, its phase equals
378
+ the common phase theta[n] = 2*pi*delta_f*n/f_s + phi_PN[n]
379
+ + phi_0 - the carrier **frequency offset and phase noise jointly**. No
380
+ symbol decisions are required, so there is no M-th-power noise enhancement
381
+ and the estimate tracks fast phase noise sample-by-sample.
382
+
383
+ The tone is isolated with a **zero-phase** spectral window (FFT -> window ->
384
+ IFFT), so there is no group-delay misalignment between the recovered phase
385
+ and the samples.
386
+
387
+ Note: operates on the **oversampled waveform, before matched filtering and
388
+ decimation**. The tone lives in a guard band that the matched filter would
389
+ otherwise remove. Apply the returned phase to the same oversampled
390
+ ``samples`` with ``correct_carrier_phase``, then run matched filtering /
391
+ decimation and any residual 1-sps CPR.
392
+
393
+ Parameters
394
+ ----------
395
+ samples : array_like
396
+ Oversampled complex samples (``sps > 1``). Shape: ``(N,)`` or
397
+ ``(C, N)``. Same rate as used for ``add_pilot_tone``.
398
+ sampling_rate : float
399
+ Sampling rate f_s in Hz.
400
+ tone_frequency : float
401
+ Nominal pilot-tone frequency f_p in Hz (as added at the TX).
402
+ The recovered phase is referenced to **this** carrier, so any carrier
403
+ frequency offset remains in the phase ramp when
404
+ ``remove_frequency_offset=False`` is *not* set (see below).
405
+ bandwidth : float
406
+ Half-width B of the spectral extraction window in Hz - the
407
+ **tracking bandwidth**. Must be wide enough to pass the phase-noise
408
+ sidebands (``B ≳ a few x linewidth``) yet narrow enough to reject the
409
+ data band (``B`` smaller than the tone-to-signal-edge guard). See the
410
+ guide at the end of this docstring.
411
+ search_band : float, optional
412
+ Half-width in Hz of the peak-search window handed to
413
+ ``find_bias_tone`` when ``refine_tone=True``.
414
+ The actual tone peak is sought within
415
+ ``[f_p - search_band, f_p + search_band]``; this bounds how far
416
+ a frequency offset may have dragged the tone from nominal. Defaults
417
+ to ``bandwidth``. Enlarge it (independently of ``bandwidth``)
418
+ when the offset can exceed ``B`` but keep it inside the guard so the
419
+ data band never wins the argmax.
420
+ refine_tone : bool, default True
421
+ If ``True``, locate the actual per-channel tone frequency with
422
+ ``find_bias_tone`` and centre the extraction
423
+ window there. Essential when a frequency offset may shift the tone by
424
+ more than ``B`` (otherwise the tone falls outside a window centred at
425
+ nominal). If ``False``, the window is centred at ``tone_frequency``.
426
+ window : str or tuple, default 'tukey'
427
+ Spectral window applied over the passband |f - f_centre| <= B.
428
+ Any spec accepted by ``get_window`` (e.g. ``'tukey'``,
429
+ ``'boxcar'``, ``('gaussian', std)``). Tukey (default) gives a flat top
430
+ with tapered edges for suppressed ringing.
431
+ remove_frequency_offset : bool, default True
432
+ If ``True`` (default), the recovered phase **retains** the linear ramp
433
+ from any residual carrier frequency offset, so applying it corrects
434
+ frequency offset and phase noise together. If ``False``, the
435
+ least-squares linear trend is subtracted per channel, leaving only the
436
+ phase-noise fluctuation (use when the frequency offset is handled by a
437
+ separate stage).
438
+ joint_channels : bool, default False
439
+ For MIMO inputs (C > 1): if ``True``, coherently sum the extracted
440
+ tone phasors across channels before taking the angle (shared-LO,
441
+ ~√C variance reduction). The single trajectory is broadcast to all
442
+ rows. No effect for SISO.
443
+ debug_plot : bool, default False
444
+ If ``True``, open the dedicated diagnostic figure
445
+ (``pilot_tone_phase_estimate``): the tone
446
+ spectrum with the extraction window overlaid, and the recovered phase.
447
+
448
+ Returns
449
+ -------
450
+ array_like
451
+ Per-sample phase estimate theta_hat[n] in radians. Shape
452
+ matches ``samples``; same backend. Apply with
453
+ ``correct_carrier_phase``.
454
+
455
+ Notes
456
+ -----
457
+ Pipeline: FFT -> (optional) refine tone centre -> zero-phase window extraction
458
+ -> strip nominal carrier -> unwrap(angle) in float64.
459
+
460
+ ``bandwidth`` B trades phase-noise tracking bandwidth against tone SNR.
461
+ Lower bound: B ≳ 3-5 * linewidth (pass all phase-noise sidebands).
462
+ Upper bound: B below the guard between the tone and the signal band edge.
463
+ Place the tone at |f_p| > (1+beta)*R_s/2 + B and keep |f_p| + B < f_s/2.
464
+ """
465
+ if bandwidth <= 0.0:
466
+ raise ValueError(f"bandwidth must be > 0, got {bandwidth}.")
467
+ if not (-sampling_rate / 2.0 < tone_frequency < sampling_rate / 2.0):
468
+ raise ValueError(f"tone_frequency={tone_frequency} must lie in (-fs/2, fs/2).")
469
+
470
+ samples, xp, _ = dispatch(samples)
471
+ was_1d = samples.ndim == 1
472
+ if was_1d:
473
+ samples = samples[None, :] # (1, N)
474
+ C, N = samples.shape
475
+
476
+ df = sampling_rate / N
477
+ if bandwidth < df:
478
+ logger.warning(
479
+ "CPR (pilot-tone): bandwidth=%.3g Hz is below the FFT resolution df=fs/N=%.3g Hz; the extraction window may capture too few bins. Increase bandwidth or the record length N.",
480
+ bandwidth,
481
+ df,
482
+ )
483
+
484
+ # Isolate the tone and strip the nominal carrier (shared core).
485
+ phasor, f_centers, _, _, W, X = _extract_pilot_phasor(
486
+ samples,
487
+ sampling_rate,
488
+ tone_frequency,
489
+ bandwidth,
490
+ xp,
491
+ search_band=search_band,
492
+ refine_tone=refine_tone,
493
+ window=window,
494
+ return_window=debug_plot,
495
+ )
496
+ n = xp.arange(N, dtype=xp.float64) # for the residual-FOE detrend below
497
+
498
+ # 5) Phase extraction + unwrap in float64.
499
+ if joint_channels and C > 1:
500
+ z_joint = xp.sum(phasor, axis=0) # (N,) coherent sum
501
+ theta_joint = xp.unwrap(xp.angle(z_joint).astype(xp.float64)) # (N,)
502
+ theta = xp.broadcast_to(theta_joint[None, :], (C, N)).copy()
503
+ else:
504
+ theta = xp.unwrap(xp.angle(phasor).astype(xp.float64), axis=-1) # (C, N)
505
+
506
+ if not remove_frequency_offset:
507
+ # Subtract the per-channel least-squares linear trend (residual FOE),
508
+ # preserving the mean phase; leaves only the phase-noise fluctuation.
509
+ nc = n - xp.mean(n)
510
+ denom = xp.sum(nc * nc)
511
+ theta_c = theta - xp.mean(theta, axis=-1, keepdims=True)
512
+ slope = xp.sum(theta_c * nc[None, :], axis=-1, keepdims=True) / denom # (C, 1)
513
+ theta = theta - slope * nc[None, :]
514
+
515
+ # Host copy of theta is needed only for the INFO summary and the optional
516
+ # debug plot; skip the transfer + reductions otherwise (the device theta is
517
+ # what gets returned and applied).
518
+ _want_log = logger.isEnabledFor(logging.INFO)
519
+ if _want_log or debug_plot:
520
+ theta_np = to_device(theta, "cpu")
521
+ if _want_log:
522
+ phi_mean_deg = float(np.mean(theta_np)) * 180.0 / np.pi
523
+ phi_std_deg = float(np.std(theta_np)) * 180.0 / np.pi
524
+ mode_str = "joint" if (joint_channels and C > 1) else "independent"
525
+ logger.info(
526
+ "CPR (pilot-tone, %s, %s): phase mean=%.2f°, std=%.2f° [f_p=%.3g Hz, B=%.3g Hz, refine=%s, remove_foe=%s, C=%s]",
527
+ window,
528
+ mode_str,
529
+ phi_mean_deg,
530
+ phi_std_deg,
531
+ tone_frequency,
532
+ bandwidth,
533
+ refine_tone,
534
+ remove_frequency_offset,
535
+ C,
536
+ )
537
+
538
+ if debug_plot:
539
+ from .. import plotting as _plotting
540
+
541
+ _plotting.plot_pilot_tone_phase_estimate(
542
+ freqs=np.fft.fftfreq(N, d=1.0 / sampling_rate),
543
+ mag_spectrum=to_device(xp.abs(X), "cpu"),
544
+ window=to_device(W, "cpu"),
545
+ f_tones=f_centers,
546
+ theta=theta_np,
547
+ tone_frequency=float(tone_frequency),
548
+ bandwidth=float(bandwidth),
549
+ show=True,
550
+ )
551
+
552
+ if was_1d:
553
+ return theta[0]
554
+ return theta
555
+
556
+
557
+ def _lowpass_fft(z: ArrayType, sampling_rate: float, cutoff: float, xp) -> ArrayType:
558
+ """Zero-phase brick-wall low-pass of a complex stream (FFT -> mask -> IFFT).
559
+
560
+ Used to isolate the **slow** inter-tone differential phasor; zero-phase so
561
+ the recovered ``δ(n)`` is lag-free (it is far inside the passband anyway).
562
+ """
563
+ N = z.shape[-1]
564
+ freqs = xp.fft.fftfreq(N, d=1.0 / sampling_rate)
565
+ mask = (xp.abs(freqs) <= cutoff).astype(z.real.dtype)
566
+ return xp.fft.ifft(xp.fft.fft(z, axis=-1) * mask, axis=-1)
567
+
568
+
569
+ def recover_carrier_phase_pilot_tones(
570
+ samples: ArrayType,
571
+ sampling_rate: float,
572
+ tone_frequencies,
573
+ bandwidth: float,
574
+ differential_bandwidth: float = 5e3,
575
+ search_band: float | None = None,
576
+ per_tone_channel: list | None = None,
577
+ snr_gate_db: float = 3.0,
578
+ coherence_gate: float = 0.3,
579
+ refine_tone: bool = True,
580
+ window: str | tuple = "tukey",
581
+ return_diagnostics: bool = False,
582
+ debug_plot: bool = False,
583
+ ):
584
+ r"""
585
+ Common carrier-phase recovery from two (or more) CW pilot tones via
586
+ SNR-weighted maximal-ratio combining with slow inter-tone tracking.
587
+
588
+ For a shared-laser/shared-LO dual-pol link the carrier (beat) phase phi[n]
589
+ is common-mode across both polarizations, and every pilot rides the same
590
+ phi[n]. Combining K tones lowers the residual phase noise by up to sqrt(K)
591
+ over a single tone, directly reducing the excess noise it converts into
592
+ (xi_phi ~ V_A * Var(d_phi)). The static inter-tone offset theta_k - theta_0
593
+ is constant back-to-back but drifts over fiber (SOP rotation acting on the
594
+ orthogonally-launched pilots), so it is tracked, not calibrated: the product
595
+ z_k * conj(z_0) cancels the common phi[n] exactly, leaving only the slow
596
+ differential, which a narrow low-pass isolates.
597
+
598
+ The whole combine collapses to one expression,
599
+
600
+ z_comb[n] = sum_k z_k[n] * conj(c_k[n]) / sigma_k^2,
601
+ c_k[n] = LPF( z_k[n] * conj(z_0[n]) ),
602
+
603
+ where conj(c_k) carries both the magnitude weight |A_k||A_0| (so fades are
604
+ down-weighted per-sample) and the de-rotation exp(-j*delta_k); dividing by
605
+ the additive-noise power sigma_k^2 makes it true MRC. A tone whose SNR or
606
+ differential coherence falls below the gates is dropped, so the combine
607
+ degrades gracefully to single-tone in a deep fade.
608
+
609
+ Operates on the oversampled waveform, before matched filtering, exactly like
610
+ ``recover_carrier_phase_pilot_tone``. Best run *after* the polarization
611
+ demux (each pilot isolated on its own output, ``per_tone_channel=[0, 1]``);
612
+ pre-demux it falls back to a joint-across-channels sum per tone.
613
+
614
+ Parameters
615
+ ----------
616
+ samples : (N,) or (C, N) array
617
+ Oversampled complex samples (``sps > 1``).
618
+ sampling_rate : float
619
+ Sampling rate in Hz (of *these* samples - pass the post-resample rate if
620
+ the demux/resample ran first).
621
+ tone_frequencies : sequence of float
622
+ Nominal pilot-tone frequencies in Hz (length ``K``).
623
+ bandwidth : float
624
+ Half-width of the per-tone extraction window in Hz (the common-phase
625
+ tracking bandwidth); see ``recover_carrier_phase_pilot_tone``.
626
+ differential_bandwidth : float, default 5e3
627
+ Low-pass cut-off in Hz for the slow inter-tone differential delta_k[n].
628
+ Choose above the SOP drift rate (so delta is not lagged) and far below
629
+ the phase-noise band (kHz-scale is typical). Set it from the knee of the
630
+ ``angle(z_k·conj(z_0))`` spectrum.
631
+ search_band : float, optional
632
+ Peak-search half-width handed to ``find_bias_tone``; defaults to
633
+ ``bandwidth``.
634
+ per_tone_channel : list of int, optional
635
+ Channel each tone is read from (post-demux isolation, e.g. ``[0, 1]``).
636
+ If ``None``, each tone's phasor is the coherent sum across all channels
637
+ (pre-demux joint combine).
638
+ snr_gate_db : float, default 3.0
639
+ A non-reference tone is dropped if its in-band SNR is below this.
640
+ coherence_gate : float, default 0.3
641
+ A non-reference tone is dropped if its differential coherence
642
+ ``mean|c_k| / sqrt(S_k·S_ref)`` is below this (deep fade / lost lock).
643
+ refine_tone, window : see ``recover_carrier_phase_pilot_tone``.
644
+ return_diagnostics : bool, default False
645
+ If ``True``, also return a dict with ``delta`` (per-tone delta_k[n]),
646
+ ``snr_db``, ``ref`` (reference-tone index) and ``used`` (combined tone
647
+ indices).
648
+ debug_plot : bool, default False
649
+ If ``True``, plot the per-tone differential phase and the combined track.
650
+
651
+ Returns
652
+ -------
653
+ array_like
654
+ Per-sample common phase estimate phi_hat[n], shape matching ``samples``
655
+ (one track broadcast to all rows). Apply with
656
+ ``correct_carrier_phase``. If ``return_diagnostics``, returns
657
+ ``(phi, diagnostics)``.
658
+ """
659
+ if bandwidth <= 0.0:
660
+ raise ValueError(f"bandwidth must be > 0, got {bandwidth}.")
661
+ tone_frequencies = list(tone_frequencies)
662
+ K = len(tone_frequencies)
663
+ if K < 1:
664
+ raise ValueError("tone_frequencies must contain at least one frequency.")
665
+
666
+ samples, xp, _ = dispatch(samples)
667
+ was_1d = samples.ndim == 1
668
+ if was_1d:
669
+ samples = samples[None, :]
670
+ C, N = samples.shape
671
+ if per_tone_channel is not None and len(per_tone_channel) != K:
672
+ raise ValueError(
673
+ f"per_tone_channel must have one entry per tone (len {K}), "
674
+ f"got {len(per_tone_channel)}."
675
+ )
676
+
677
+ # 1) Extract each tone's scalar phasor stream z_k(n) and its (S_k, σ_k²).
678
+ # One shared working-precision FFT of the record serves every tone's
679
+ # refinement and extraction (K windowed IFFTs remain, nothing else scales
680
+ # with K·N).
681
+ xw = (
682
+ samples
683
+ if samples.dtype == xp.complex128
684
+ else samples.astype(xp.complex64, copy=False)
685
+ )
686
+ X = xp.fft.fft(xw, axis=-1) # (C, N)
687
+ z_tones, sig, noise, f_centers = [], [], [], []
688
+ for k, f_k in enumerate(tone_frequencies):
689
+ ph, fc, s_c, n_c, _, _ = _extract_pilot_phasor(
690
+ samples,
691
+ sampling_rate,
692
+ f_k,
693
+ bandwidth,
694
+ xp,
695
+ search_band=search_band,
696
+ refine_tone=refine_tone,
697
+ window=window,
698
+ X=X,
699
+ )
700
+ if per_tone_channel is None:
701
+ z_k = xp.sum(ph, axis=0) # joint across channels (pre-demux)
702
+ s_k, n_k = float(np.sum(s_c)), float(np.sum(n_c))
703
+ else:
704
+ ch = int(per_tone_channel[k])
705
+ z_k, s_k, n_k = ph[ch], float(s_c[ch]), float(n_c[ch])
706
+ z_tones.append(z_k)
707
+ sig.append(s_k)
708
+ noise.append(max(n_k, 1e-30))
709
+ f_centers.append(fc)
710
+
711
+ # 2) Reference = highest-SNR tone; build the slow differential phasors c_k.
712
+ snr = np.array([s / nz for s, nz in zip(sig, noise)], dtype=np.float64)
713
+ ref = int(np.argmax(snr))
714
+ z_ref = z_tones[ref]
715
+ snr_gate = 10.0 ** (snr_gate_db / 10.0)
716
+
717
+ z_comb = xp.zeros(N, dtype=z_ref.dtype)
718
+ delta_diag, used = [], []
719
+ want_diag = return_diagnostics or debug_plot # per-tone δ_k needed either way
720
+ for k in range(K):
721
+ if k == ref:
722
+ # Self-product: LPF(|z|²) ≈ |A|² + σ²; subtract the floor so the
723
+ # reference weight is the true |A|² (its phase is 0 => no de-rotation).
724
+ c_k = _lowpass_fft(
725
+ (xp.abs(z_ref) ** 2).astype(z_ref.dtype),
726
+ sampling_rate,
727
+ differential_bandwidth,
728
+ xp,
729
+ )
730
+ c_k = xp.clip(xp.real(c_k) - noise[ref], 1e-30, None).astype(z_ref.dtype)
731
+ coh = 1.0
732
+ else:
733
+ # Cross-product: the two tones' noises are independent => the LPF
734
+ # rejects them, so c_k ≈ A_k A_0* - carries |A_k||A_0| and e^{jδ_k}.
735
+ c_k = _lowpass_fft(
736
+ z_tones[k] * xp.conj(z_ref),
737
+ sampling_rate,
738
+ differential_bandwidth,
739
+ xp,
740
+ )
741
+ coh = float(xp.mean(xp.abs(c_k))) / np.sqrt(max(sig[k] * sig[ref], 1e-30))
742
+ if snr[k] < snr_gate or coh < coherence_gate:
743
+ logger.info(
744
+ "CPR (pilot-tones): tone %s dropped (SNR=%.1f dB, coherence=%.2f).",
745
+ k,
746
+ 10 * np.log10(snr[k]),
747
+ coh,
748
+ )
749
+ if want_diag:
750
+ delta_diag.append(
751
+ to_device(xp.angle(c_k).astype(xp.float64), "cpu")
752
+ )
753
+ continue
754
+ contrib = z_tones[k] * xp.conj(c_k)
755
+ contrib /= noise[k]
756
+ z_comb += contrib
757
+ used.append(k)
758
+ if want_diag:
759
+ delta_diag.append(to_device(xp.angle(c_k).astype(xp.float64), "cpu"))
760
+
761
+ # 3) Common phase = angle of the combined phasor, unwrapped in float64.
762
+ phi = xp.unwrap(xp.angle(z_comb).astype(xp.float64)) # (N,)
763
+ phi_full = xp.broadcast_to(phi[None, :], (C, N)).copy()
764
+
765
+ # Host copy of phi is needed only for the INFO summary and the optional
766
+ # debug plot; skip the transfer otherwise (phi_full drives the correction).
767
+ _want_log = logger.isEnabledFor(logging.INFO)
768
+ if _want_log or debug_plot:
769
+ phi_np = to_device(phi, "cpu")
770
+ if _want_log:
771
+ logger.info(
772
+ "CPR (pilot-tones, MRC): phase std=%.2f°, [K=%s, used=%s, ref=%s, B=%.3g Hz, diff_B=%.3g Hz, C=%s]",
773
+ float(np.std(phi_np)) * 180 / np.pi,
774
+ K,
775
+ used,
776
+ ref,
777
+ bandwidth,
778
+ differential_bandwidth,
779
+ C,
780
+ )
781
+
782
+ if debug_plot:
783
+ from .. import plotting as _plotting
784
+
785
+ _plotting.plot_pilot_tones_phase_estimate(
786
+ delta=delta_diag,
787
+ phi=phi_np,
788
+ ref=ref,
789
+ used=used,
790
+ show=True,
791
+ )
792
+
793
+ phi_out = phi_full[0] if was_1d else phi_full
794
+ if return_diagnostics:
795
+ diagnostics = {
796
+ "delta": delta_diag,
797
+ "snr_db": 10.0 * np.log10(snr),
798
+ "ref": ref,
799
+ "used": used,
800
+ "f_centers": [np.asarray(fc) for fc in f_centers],
801
+ }
802
+ return phi_out, diagnostics
803
+ return phi_out