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,751 @@
1
+ """Phase corrections, cycle-slip repair, and ambiguity resolution."""
2
+
3
+ import logging
4
+
5
+ import numpy as np
6
+
7
+ from ..backend import ArrayType, dispatch, to_device
8
+ from ..core.signal import Signal
9
+ from ..logger import logger
10
+
11
+
12
+ def smooth_phase_wiener(
13
+ phase: ArrayType,
14
+ process_variance: float | None = None,
15
+ measurement_variance: float | None = None,
16
+ linewidth: float | None = None,
17
+ sampling_rate: float | None = None,
18
+ detrend: bool = True,
19
+ ) -> ArrayType:
20
+ r"""
21
+ Zero-phase Wiener smoother for a random-walk (Wiener) carrier phase.
22
+
23
+ Optimal minimum-variance estimate of a phase that random-walks with
24
+ per-sample increment variance ``q`` (the random-walk strength) observed in
25
+ white phase-estimation noise of variance ``r``. Applies the non-causal
26
+ Wiener filter
27
+
28
+ H(w) = S_phi(w) / (S_phi(w) + r),
29
+ S_phi(w) = q / (2 - 2*cos(w)),
30
+
31
+ in the frequency domain (FFT -> multiply by the real, even H -> IFFT), so it
32
+ is zero-phase (no group delay) and O(N log N). This is the principled way to
33
+ hit the smallest residual phase std for a given ``q`` and ``r`` - it trades
34
+ tracking lag against additive noise automatically, where a fixed extraction
35
+ bandwidth must be tuned by hand.
36
+
37
+ The smoother is agnostic to where the track came from: it needs only the two
38
+ scalars ``q`` and ``r``. ``q`` may be given directly or derived from a
39
+ linewidth (see ``process_variance`` / ``linewidth`` below); ``r`` is supplied
40
+ as ``measurement_variance``, however the caller measured it.
41
+
42
+ Because it only rescales the phase track (a deterministic, sample-independent
43
+ low-pass), it stays a unit-modulus correction downstream and cannot hide
44
+ excess noise when applied to a common reference phase rather than the data
45
+ samples.
46
+
47
+ Parameters
48
+ ----------
49
+ phase : (N,) or (C, N) array
50
+ Unwrapped phase track (e.g. from a ``recover_carrier_phase_pilot_tone*``
51
+ function), in radians.
52
+ process_variance : float, optional
53
+ Per-sample phase-increment variance q [rad²] - the random-walk strength.
54
+ Provide this, or derive it from ``linewidth`` + ``sampling_rate`` (see
55
+ below). Exactly one of the two routes is required.
56
+ measurement_variance : float, optional
57
+ Phase-estimation noise variance r [rad²] - the per-sample variance of the
58
+ additive phase-measurement noise, however it was measured. Required.
59
+ linewidth, sampling_rate : float, optional
60
+ Convenience route to ``process_variance``: q = 2*pi*linewidth / f_s, where
61
+ ``linewidth`` is the combined oscillator linewidth [Hz] and
62
+ ``sampling_rate`` is f_s [Hz] of the ``phase`` track. Both must be given
63
+ together, and only when ``process_variance`` is omitted.
64
+ detrend : bool, default True
65
+ Remove the per-channel mean + linear trend before filtering and add it
66
+ back after. Recommended: the random-walk PSD diverges at DC, so a raw
67
+ ramp (residual frequency offset) would be distorted; detrending keeps it
68
+ exact.
69
+
70
+ Returns
71
+ -------
72
+ array_like
73
+ Smoothed phase, same shape and backend as ``phase``.
74
+ """
75
+ if process_variance is None:
76
+ if linewidth is None or sampling_rate is None:
77
+ raise ValueError(
78
+ "Provide process_variance, or both linewidth and sampling_rate."
79
+ )
80
+ process_variance = 2.0 * np.pi * float(linewidth) / float(sampling_rate)
81
+ if measurement_variance is None:
82
+ raise ValueError("Provide measurement_variance.")
83
+ q, r = float(process_variance), float(measurement_variance)
84
+ if q <= 0.0 or r <= 0.0:
85
+ raise ValueError(f"process/measurement variance must be > 0, got q={q}, r={r}.")
86
+
87
+ phase, xp, _ = dispatch(phase)
88
+ was_1d = phase.ndim == 1
89
+ if was_1d:
90
+ phase = phase[None, :]
91
+ C, N = phase.shape
92
+ phi = phase.astype(xp.float64)
93
+
94
+ # Detrend per channel (mean + linear) so the DC-divergent random-walk PSD
95
+ # does not distort the residual-FOE ramp; restore the trend after filtering.
96
+ n = xp.arange(N, dtype=xp.float64)
97
+ if detrend:
98
+ nc = n - xp.mean(n)
99
+ denom = xp.sum(nc * nc)
100
+ mean = xp.mean(phi, axis=-1, keepdims=True)
101
+ slope = xp.sum((phi - mean) * nc[None, :], axis=-1, keepdims=True) / denom
102
+ trend = mean + slope * nc[None, :]
103
+ else:
104
+ trend = xp.zeros((C, 1), dtype=xp.float64)
105
+ phi_c = phi - trend
106
+
107
+ # Real, even Wiener gain H(ω); keep DC (H[0]=1) where S_φ -> ∞. The phase
108
+ # track is real, so the filter runs on the half spectrum (rfft/irfft) -
109
+ # half the transform work and spectrum memory of the full complex FFT.
110
+ omega = 2.0 * xp.pi * xp.fft.rfftfreq(N)
111
+ denom_w = 2.0 - 2.0 * xp.cos(omega)
112
+ denom_w = xp.where(denom_w <= 0.0, xp.full_like(denom_w, 1e-300), denom_w)
113
+ S = q / denom_w
114
+ H = S / (S + r)
115
+ H[0] = 1.0
116
+
117
+ phi_s = xp.fft.irfft(xp.fft.rfft(phi_c, axis=-1) * H[None, :], n=N, axis=-1)
118
+ phi_s = phi_s + trend
119
+
120
+ if logger.isEnabledFor(logging.INFO):
121
+ # Two std reductions + host syncs, needed only for the line below.
122
+ std_in = float(xp.std(phi_c))
123
+ std_out = float(xp.std(phi_s - trend))
124
+ logger.info(
125
+ "Wiener phase smoother: q=%.3g, r=%.3g rad², residual std %.2f° -> %.2f°.",
126
+ q,
127
+ r,
128
+ np.degrees(std_in),
129
+ np.degrees(std_out),
130
+ )
131
+
132
+ return phi_s[0] if was_1d else phi_s
133
+
134
+
135
+ _PHASE_ROTATE_KERNEL: dict = {}
136
+
137
+
138
+ def _get_cupy_phase_rotate():
139
+ """Compile and cache the fused CuPy phase-rotation kernel.
140
+
141
+ Fuses the whole ``s · exp(-j·wrap(φ))`` chain - float64 wrap, float32
142
+ sin/cos, complex multiply - into a single elementwise kernel: one read of
143
+ the symbols, one read of the phase, one write, instead of the ~7 separate
144
+ full-record kernel passes and temporaries of the ufunc chain.
145
+ """
146
+ if "k" not in _PHASE_ROTATE_KERNEL:
147
+ import cupy as cp
148
+
149
+ _PHASE_ROTATE_KERNEL["k"] = cp.ElementwiseKernel(
150
+ "complex64 s, float64 phi",
151
+ "complex64 out",
152
+ """
153
+ double w = phi - rint(phi * 0.15915494309189535) * 6.283185307179586;
154
+ float sw, cw;
155
+ sincosf((float)w, &sw, &cw);
156
+ out = s * complex<float>(cw, -sw);
157
+ """,
158
+ "commkit_phase_rotate",
159
+ )
160
+ return _PHASE_ROTATE_KERNEL["k"]
161
+
162
+
163
+ def correct_carrier_phase(
164
+ symbols: ArrayType,
165
+ phase_vector: ArrayType,
166
+ ) -> ArrayType:
167
+ """
168
+ Applies carrier phase correction to a symbol sequence.
169
+
170
+ Rotates each symbol by the negative of the estimated phase to cancel
171
+ the carrier phase offset: y[n] = s[n] * exp(-j * phi_hat[n]).
172
+
173
+ Parameters
174
+ ----------
175
+ symbols : array_like
176
+ Complex symbols. Shape: (N,) or (C, N).
177
+ phase_vector : array_like
178
+ Per-symbol phase estimates in radians. Shape: (N,) for SISO, or
179
+ broadcastable to ``symbols.shape`` for MIMO.
180
+
181
+ Returns
182
+ -------
183
+ array_like
184
+ Phase-corrected symbols, same shape and dtype as ``symbols``.
185
+ """
186
+ symbols, xp, _ = dispatch(symbols)
187
+ logger.debug("Applying carrier phase correction: shape=%s", symbols.shape)
188
+ # Wrap to [-π, π] in float64 (handles unbounded phase trajectories from
189
+ # standalone CPR), then rotate with a float32 phasor.
190
+ phase_f64 = xp.asarray(phase_vector, dtype=xp.float64)
191
+ if xp is not np and symbols.dtype == xp.complex64:
192
+ # GPU fast path: single fused kernel (broadcasts (N,) phase over (C, N)).
193
+ return _get_cupy_phase_rotate()(symbols, phase_f64)
194
+ two_pi = 2.0 * xp.pi
195
+ phase_wrapped = (phase_f64 - xp.round(phase_f64 / two_pi) * two_pi).astype(
196
+ xp.float32
197
+ )
198
+ phasor = xp.exp(-1j * phase_wrapped)
199
+ if phasor.dtype != symbols.dtype:
200
+ phasor = phasor.astype(symbols.dtype)
201
+ return symbols * phasor
202
+
203
+
204
+ _NUMBA_CYCLE_SLIP: dict = {}
205
+
206
+
207
+ def _get_numba_cycle_slip():
208
+ """JIT-compile and cache the Numba cycle-slip correction kernel.
209
+
210
+ Returns
211
+ -------
212
+ callable
213
+ Numba-compiled ``_cycle_slip_loop``.
214
+ """
215
+ if "cs" not in _NUMBA_CYCLE_SLIP:
216
+ import numba
217
+
218
+ @numba.njit(cache=True, fastmath=True, nogil=True)
219
+ def _cycle_slip_loop(phi_u, symmetry, history_length, threshold):
220
+ """Cycle-slip detection and correction via linear extrapolation.
221
+
222
+ Scans the block-phase trajectory ``phi_u`` sequentially. For each
223
+ block, linearly extrapolates from up to ``history_length`` past
224
+ *corrected* blocks. When the deviation exceeds ``threshold``, a
225
+ ``π/2`` step correction is applied.
226
+
227
+ The linear regression uses relative coordinates [0, W-1] so that
228
+ ``Sx`` and ``Sxx`` are exact compile-time constants. Only ``Sy``
229
+ and ``Sxy`` are maintained as running state, updated in O(1) per
230
+ step via a closed-form sliding-window identity.
231
+
232
+ Parameters
233
+ ----------
234
+ phi_u : (B,) float64
235
+ Block-phase trajectory after M-fold unwrap (modified in place).
236
+ symmetry : int
237
+ Rotational symmetry order; correction quantum = ``2π/symmetry``.
238
+ Pass 4 for QAM (all BPS, VV, Tikhonov use 4-fold symmetry).
239
+ history_length : int
240
+ Maximum number of past corrected blocks used for extrapolation.
241
+ Use ``min(b, history_length)`` at each step.
242
+ threshold : float64
243
+ Deviation from extrapolated value that triggers a correction
244
+ (radians). Default in the caller: ``π/4``.
245
+
246
+ Returns
247
+ -------
248
+ (B,) float64
249
+ Corrected block-phase trajectory (same array, modified in place).
250
+ """
251
+ two_pi = 2.0 * np.pi
252
+ quantum = two_pi / float(symmetry)
253
+ B = len(phi_u)
254
+ W = history_length
255
+ W_f = float(W)
256
+
257
+ # Precompute full-window regression constants in relative coords [0, W-1].
258
+ # With relative coords the x-values are always small integers, so Sx and
259
+ # Sxx never grow and there is no catastrophic cancellation regardless of
260
+ # how many total blocks have been processed.
261
+ Sx_full = W_f * (W_f - 1.0) / 2.0
262
+ Sxx_full = W_f * (W_f - 1.0) * (2.0 * W_f - 1.0) / 6.0
263
+ denom_full = W_f * Sxx_full - Sx_full * Sx_full # W²(W²-1)/12
264
+
265
+ # Only Sy and Sxy need to be maintained as running state.
266
+ buf_y = np.empty(W, dtype=np.float64)
267
+ buf_head = 0 # next write slot (circular)
268
+ n_buf = 0 # valid entries currently in buffer
269
+
270
+ Sy = 0.0
271
+ Sxy = 0.0
272
+
273
+ for b in range(B):
274
+ y_b = phi_u[b]
275
+
276
+ if n_buf == 0:
277
+ # First block: trust it unconditionally (at relative position 0).
278
+ buf_y[0] = y_b
279
+ buf_head = 1
280
+ n_buf = 1
281
+ Sy = y_b
282
+ Sxy = 0.0 # 0 * y_b
283
+ continue
284
+
285
+ if n_buf < min(10, W):
286
+ # Constant extrapolation during warmup to avoid cementing false slips.
287
+ phi_pred = buf_y[(buf_head - 1) % W]
288
+ else:
289
+ # Linear extrapolation. Prediction target is always one step past
290
+ # the newest buffered entry, i.e. relative coordinate = n_buf.
291
+ x_pred = float(n_buf)
292
+ n_f = float(n_buf)
293
+ if n_buf < W:
294
+ # Partial window: derive exact Sx/Sxx from closed-form sums.
295
+ Sx_p = n_f * (n_f - 1.0) / 2.0
296
+ Sxx_p = n_f * (n_f - 1.0) * (2.0 * n_f - 1.0) / 6.0
297
+ denom = n_f * Sxx_p - Sx_p * Sx_p
298
+ if abs(denom) > 1e-30:
299
+ slope = (n_f * Sxy - Sx_p * Sy) / denom
300
+ intercept = (Sy - slope * Sx_p) / n_f
301
+ else:
302
+ slope = 0.0
303
+ intercept = Sy / n_f
304
+ else:
305
+ # Full window: use precomputed constants (numerically exact).
306
+ if denom_full > 1e-30:
307
+ slope = (W_f * Sxy - Sx_full * Sy) / denom_full
308
+ intercept = (Sy - slope * Sx_full) / W_f
309
+ else:
310
+ slope = 0.0
311
+ intercept = Sy / W_f
312
+ phi_pred = slope * x_pred + intercept
313
+
314
+ diff = y_b - phi_pred
315
+ # Round to nearest correction quantum
316
+ k = round(diff / quantum)
317
+ if abs(diff) > threshold and k != 0:
318
+ phi_u[b] -= float(k) * quantum
319
+ y_b = phi_u[b]
320
+
321
+ # Update circular buffer using relative coordinates.
322
+ if n_buf == W:
323
+ # Slide window: evict oldest (relative pos 0), shift all down by 1,
324
+ # add y_b at relative position W-1.
325
+ # Sxy update uses the identity:
326
+ # Sxy_new = Sxy_old - Sy_old + y_old + (W-1)·y_new
327
+ # (derived by relabelling positions after eviction)
328
+ old_idx = buf_head % W
329
+ y_old = buf_y[old_idx]
330
+ Sxy = Sxy - Sy + y_old + (W_f - 1.0) * y_b # must precede Sy update
331
+ Sy = Sy - y_old + y_b
332
+ buf_y[old_idx] = y_b
333
+ buf_head += 1
334
+ else:
335
+ # Append at relative position n_buf.
336
+ idx = buf_head % W
337
+ buf_y[idx] = y_b
338
+ Sxy += float(n_buf) * y_b
339
+ Sy += y_b
340
+ buf_head += 1
341
+ n_buf += 1
342
+
343
+ return phi_u
344
+
345
+ _NUMBA_CYCLE_SLIP["cs"] = _cycle_slip_loop
346
+
347
+ return _NUMBA_CYCLE_SLIP["cs"]
348
+
349
+
350
+ def correct_cycle_slips(
351
+ phi_u: np.ndarray,
352
+ symmetry: int = 4,
353
+ history_length: int = 1000,
354
+ threshold: float = np.pi / 4,
355
+ ) -> np.ndarray:
356
+ """
357
+ Detects and corrects cycle slips in a block-phase trajectory.
358
+
359
+ After ``xp.unwrap`` resolves the M-fold ambiguity, residual cycle slips
360
+ may remain where the unwrapper chose the wrong quadrant. This function
361
+ scans the trajectory sequentially: for each block it extrapolates the
362
+ expected phase from up to ``history_length`` past corrected blocks using
363
+ a linear fit. When the deviation exceeds ``threshold``, the block is
364
+ corrected by the nearest integer multiple of ``2π/symmetry``.
365
+
366
+ Algorithm: linear extrapolation from the previous
367
+ ``history_length`` corrected phases; correction quantum = ``π/2`` for
368
+ 4-fold QAM symmetry; threshold = ``π/4``.
369
+
370
+ Parameters
371
+ ----------
372
+ phi_u : (B,) float64
373
+ Block-phase trajectory on CPU after M-fold unwrap (e.g. output of
374
+ ``xp.unwrap(phi_raw * M) / M``). **Modified in place.**
375
+ symmetry : int, default 4
376
+ Rotational symmetry order of the constellation. Correction quantum
377
+ is ``2π/symmetry``. Use 4 for all square QAM constellations and BPS
378
+ (which always searches over ``[0, π/2)``). For M-PSK use ``symmetry = M``.
379
+ history_length : int, default 1000
380
+ Number of past corrected blocks used for linear extrapolation.
381
+ Reduce for short bursts.
382
+ threshold : float, default π/4
383
+ Deviation from the extrapolated phase that triggers a correction.
384
+ ``π/4`` is the midpoint between adjacent correction quanta for 4-fold
385
+ symmetry.
386
+
387
+ Returns
388
+ -------
389
+ (B,) float64
390
+ Corrected block-phase trajectory (same NumPy array).
391
+
392
+ Notes
393
+ -----
394
+ Runs on CPU only (sequential scan; Numba-compiled).
395
+ The caller should transfer ``phi_u`` to CPU before calling and move
396
+ the result back to the device if needed.
397
+ """
398
+ phi_u = np.asarray(phi_u, dtype=np.float64)
399
+ kernel = _get_numba_cycle_slip()
400
+ return kernel(phi_u, int(symmetry), int(history_length), float(threshold))
401
+
402
+
403
+ def resolve_channel_permutation(
404
+ symbols: ArrayType | Signal,
405
+ ref_symbols: ArrayType | None = None,
406
+ *,
407
+ num_skip_symbols: int = 0,
408
+ ) -> ArrayType | Signal:
409
+ """Resolve a polarization (channel) permutation after MIMO equalization.
410
+
411
+ A MIMO (butterfly) equalizer has a **polarization-permutation ambiguity**:
412
+ it may emit the streams in swapped output order (output 0 carries pol 1,
413
+ etc.) - a perfectly valid demux that per-channel metrics would otherwise
414
+ score as random, since they compare ``output[i]`` with ``ref[i]``. This
415
+ matches each output stream to the reference stream it actually carries (the
416
+ bijective assignment maximizing the **rotation-invariant** cross-correlation
417
+ magnitude ``|Σ yᵢ · conj(sⱼ)|``) and reorders ``symbols`` to ``ref_symbols``
418
+ order.
419
+
420
+ Run this **before** ``resolve_phase_ambiguity`` (it is rotation
421
+ invariant, so the two compose) and before SER/BER. For a converged
422
+ *data-aided* equalizer the outputs are already pinned to the training order,
423
+ so this is a no-op; it is the robust fix for **blind** equalizers, whose
424
+ output order is arbitrary. Only a *constant* permutation is resolved - a
425
+ mid-stream swap is an equalizer-tracking issue, not a labeling one.
426
+
427
+ Parameters
428
+ ----------
429
+ symbols : array_like
430
+ Recovered symbols, ``(N,)`` or ``(C, N)``. Returned unchanged for SISO.
431
+ ref_symbols : array_like
432
+ Known transmitted symbols, same layout as ``symbols`` (the full
433
+ sequence, a pilot subset, or any known reference).
434
+ num_skip_symbols : int, default 0
435
+ Leading symbols excluded from the correlation scoring (e.g. an
436
+ unconverged transient). The reorder still covers the full input.
437
+
438
+ Returns
439
+ -------
440
+ array_like
441
+ ``symbols`` with channels reordered to match ``ref_symbols``; same
442
+ shape, dtype, and backend.
443
+ When ``symbols`` is a :class:`Signal`, ``resolved_symbols`` is reordered
444
+ against ``source_symbols`` and a new :class:`Signal` is returned.
445
+ """
446
+ if isinstance(symbols, Signal):
447
+ sig = symbols
448
+ if sig.resolved_symbols is None:
449
+ raise ValueError(
450
+ "resolved_symbols is not set. Call resolve_symbols(sig) or assign "
451
+ "resolved_symbols before calling resolve_channel_permutation()."
452
+ )
453
+ if sig.source_symbols is None:
454
+ raise ValueError(
455
+ "source_symbols is not set. Populate source_symbols (the known TX "
456
+ "symbol sequence) before calling resolve_channel_permutation()."
457
+ )
458
+ new = sig.copy()
459
+ new.resolved_symbols = resolve_channel_permutation(
460
+ sig.resolved_symbols,
461
+ sig.source_symbols,
462
+ num_skip_symbols=num_skip_symbols,
463
+ )
464
+ return new
465
+
466
+ if ref_symbols is None:
467
+ raise ValueError("resolve_channel_permutation() requires ref_symbols.")
468
+
469
+ from scipy.optimize import linear_sum_assignment
470
+
471
+ symbols, xp, _ = dispatch(symbols)
472
+ was_1d = symbols.ndim == 1
473
+ if was_1d:
474
+ return symbols
475
+ C, N = symbols.shape
476
+ if C == 1:
477
+ return symbols
478
+
479
+ ref = xp.asarray(ref_symbols)
480
+ if ref.ndim == 1:
481
+ ref = ref[None, :]
482
+ n = min(N, ref.shape[-1])
483
+ y = symbols[:, num_skip_symbols:n]
484
+ s = ref[:, num_skip_symbols:n]
485
+
486
+ # Rotation-invariant coherence matrix M[i, j] = |<y_i, s_j>| / (||y_i|| ||s_j||).
487
+ yn = y / xp.maximum(xp.linalg.norm(y, axis=-1, keepdims=True), 1e-12)
488
+ sn = s / xp.maximum(xp.linalg.norm(s, axis=-1, keepdims=True), 1e-12)
489
+ M = to_device(xp.abs(yn @ xp.conj(sn).T), "cpu") # (C_out, C_ref)
490
+
491
+ _, perm = linear_sum_assignment(-M) # perm[i] = ref stream matched by output i
492
+ perm = np.asarray(perm)
493
+ inv = np.argsort(perm) # reorder: out'[j] is the output carrying ref j
494
+
495
+ assigned = M[np.arange(C), perm]
496
+ is_identity = bool(np.array_equal(perm, np.arange(C)))
497
+ matrix_str = np.array2string(M, precision=2, suppress_small=True)
498
+ if float(assigned.min()) < 0.3:
499
+ # An output did not lock to any distinct reference stream - the demux
500
+ # likely collapsed (both outputs on one pol) rather than swapped.
501
+ logger.warning(
502
+ "resolve_channel_permutation: weak match (min coherence %.2f) - streams may not be cleanly separated (EQ collapse?). Applying best assignment %s anyway. Coherence matrix (rows=out, cols=ref):\n%s",
503
+ float(assigned.min()),
504
+ perm.tolist(),
505
+ matrix_str,
506
+ )
507
+ elif is_identity:
508
+ logger.info(
509
+ "resolve_channel_permutation: identity %s (no swap).", perm.tolist()
510
+ )
511
+ else:
512
+ logger.info(
513
+ "resolve_channel_permutation: POLARIZATION SWAP %s - reordering outputs to reference order. Coherence matrix (rows=out, cols=ref):\n%s",
514
+ perm.tolist(),
515
+ matrix_str,
516
+ )
517
+
518
+ return symbols[xp.asarray(inv)]
519
+
520
+
521
+ def resolve_phase_ambiguity(
522
+ symbols: ArrayType | Signal,
523
+ ref_symbols: ArrayType | None = None,
524
+ modulation: str | None = None,
525
+ order: int | None = None,
526
+ symmetry_order: int | None = None,
527
+ num_skip_symbols: int = 0,
528
+ pmf: np.ndarray | None = None,
529
+ ) -> ArrayType | Signal:
530
+ """
531
+ Resolves rotational phase ambiguity after blind carrier phase recovery.
532
+
533
+ Blind CPR methods (VV, BPS, Tikhonov) cannot distinguish between
534
+ ``symmetry_order`` rotational copies of the constellation. This function
535
+ tests all candidate rotations, scores each by Symbol Error Rate (SER)
536
+ against the known transmitted symbols, and returns the symbols rotated by
537
+ the best candidate.
538
+
539
+ For MIMO inputs each channel is resolved independently - after MIMO
540
+ equalisation the output streams may land on different ambiguity branches.
541
+
542
+ Parameters
543
+ ----------
544
+ symbols : array_like
545
+ Received complex symbols after CPR and ``correct_carrier_phase``.
546
+ Shape: ``(N,)`` or ``(C, N)``.
547
+ ref_symbols : array_like
548
+ Known transmitted symbols (unit-average-power normalised).
549
+ Shape: ``(N,)`` or ``(C, N)``.
550
+ modulation : str
551
+ Modulation scheme (case-insensitive): ``'qam'``, ``'psk'``, etc.
552
+ order : int
553
+ Modulation order.
554
+ symmetry_order : int, optional
555
+ Number of rotationally equivalent constellation copies to test.
556
+ Defaults to 4 for QAM (4-fold ``π/2`` symmetry) and ``order`` for
557
+ PSK. Override for non-standard constellations.
558
+ num_skip_symbols : int, default 0
559
+ Number of leading symbols to exclude from SER scoring. The applied
560
+ rotation still covers the full input - only the scoring window is
561
+ trimmed. Useful when the first ``num_skip_symbols`` symbols have not
562
+ yet converged and would bias the rotation choice. Must be strictly
563
+ less than the total symbol count.
564
+ pmf : np.ndarray, optional
565
+ Symbol PMF of shape ``(order,)`` for PS-QAM. Forwarded to
566
+ ``ser`` so the diagnostic SER reported in
567
+ the log is unbiased for shaped constellations. The phase-rotation
568
+ choice itself uses a scale-invariant inner product and does not
569
+ depend on ``pmf``.
570
+
571
+ Returns
572
+ -------
573
+ array_like
574
+ Phase-ambiguity-resolved symbols, same shape and dtype as ``symbols``.
575
+
576
+ When ``symbols`` is a :class:`Signal`, ``resolved_symbols`` is resolved
577
+ against ``source_symbols`` (using the signal's modulation/order/pmf) and a
578
+ new :class:`Signal` is returned.
579
+ """
580
+ if isinstance(symbols, Signal):
581
+ sig = symbols
582
+ if sig.resolved_symbols is None:
583
+ raise ValueError(
584
+ "resolved_symbols is not set. Call resolve_symbols(sig) or assign "
585
+ "resolved_symbols before calling resolve_phase_ambiguity()."
586
+ )
587
+ if sig.source_symbols is None:
588
+ raise ValueError(
589
+ "source_symbols is not set. Populate source_symbols (the known TX "
590
+ "symbol sequence) before calling resolve_phase_ambiguity()."
591
+ )
592
+ if sig.mod_scheme is None or sig.mod_order is None:
593
+ raise ValueError("mod_scheme and mod_order must be set.")
594
+ new = sig.copy()
595
+ new.resolved_symbols = resolve_phase_ambiguity(
596
+ sig.resolved_symbols,
597
+ sig.source_symbols,
598
+ sig.mod_scheme,
599
+ sig.mod_order,
600
+ symmetry_order=symmetry_order,
601
+ num_skip_symbols=num_skip_symbols,
602
+ pmf=sig.ps_pmf,
603
+ )
604
+ return new
605
+
606
+ if ref_symbols is None or modulation is None or order is None:
607
+ raise ValueError(
608
+ "resolve_phase_ambiguity() requires ref_symbols, modulation, and order."
609
+ )
610
+
611
+ from ..metrics import ser as _ser
612
+
613
+ symbols, xp, _ = dispatch(symbols)
614
+ was_1d = symbols.ndim == 1
615
+ if was_1d:
616
+ symbols = symbols[None, :]
617
+ C, N = symbols.shape
618
+
619
+ if num_skip_symbols >= N:
620
+ raise ValueError(
621
+ f"num_skip_symbols={num_skip_symbols} must be less than the total "
622
+ f"symbol count N={N}."
623
+ )
624
+
625
+ ref = xp.asarray(ref_symbols)
626
+ if ref.ndim == 1:
627
+ ref = ref[None, :]
628
+ if ref.shape[0] == 1 and C > 1:
629
+ ref = xp.broadcast_to(ref, (C, N))
630
+
631
+ if symmetry_order is None:
632
+ symmetry_order = 4 if "qam" in modulation.lower() else order
633
+
634
+ step = 2.0 * np.pi / symmetry_order
635
+
636
+ # ML phase ambiguity estimator: the optimal rotation maximises
637
+ # Re(e^{jkθ} · Σ y_n s_n*), which equals choosing k closest to
638
+ # -∠(Σ y_n s_n*) / step. Single inner product replaces symmetry_order
639
+ # full SER passes. All channels batched: one D2H of the (C,) angles
640
+ # instead of one float() sync per channel.
641
+ seg_y = symbols[:, num_skip_symbols:]
642
+ seg_r = ref[:, num_skip_symbols:]
643
+ corr = xp.sum(seg_y * xp.conj(seg_r), axis=-1) # (C,)
644
+ theta_np = -to_device(xp.angle(corr), "cpu") # (C,) float64, one transfer
645
+ best_k_np = np.round(theta_np / step).astype(np.int64) % symmetry_order
646
+ phasors = xp.asarray(
647
+ np.exp(1j * best_k_np * step).astype(symbols.dtype)
648
+ ) # (C,) - built on host from host indices, single H2D
649
+ out = symbols * phasors[:, None]
650
+
651
+ # SER is diagnostic-only: skip the per-channel reduction syncs entirely
652
+ # when INFO logging is disabled.
653
+ if logger.isEnabledFor(logging.INFO):
654
+ for ch in range(C):
655
+ best_ser = float(
656
+ xp.mean(
657
+ xp.asarray(
658
+ _ser(
659
+ out[ch, num_skip_symbols:],
660
+ seg_r[ch],
661
+ modulation,
662
+ order,
663
+ pmf=pmf,
664
+ )
665
+ )
666
+ )
667
+ )
668
+ logger.info(
669
+ "Phase ambiguity resolution: ch=%s, best_k=%s, rotation=%.1f°, SER=%.4f",
670
+ ch,
671
+ int(best_k_np[ch]),
672
+ best_k_np[ch] * step * 180.0 / np.pi,
673
+ best_ser,
674
+ )
675
+
676
+ if was_1d:
677
+ return out[0]
678
+ return out
679
+
680
+
681
+ def correct_phase_rotation(
682
+ symbols: ArrayType,
683
+ ref_symbols: ArrayType,
684
+ num_skip_symbols: int = 0,
685
+ ) -> ArrayType:
686
+ """Correct the static per-channel phase rotation using a reference sequence.
687
+
688
+ A rotationally-invariant blind equalizer (CMA, RDE) leaves an arbitrary
689
+ constant phase offset on each output channel - not limited to the discrete
690
+ ``k·π/M`` grid that ``resolve_phase_ambiguity`` tests. This function
691
+ estimates the continuous rotation per channel via the ML inner-product
692
+ estimator ``θ = -∠(Σ y·s*)`` over a known reference sequence and applies
693
+ the correction to the full symbol block.
694
+
695
+ The reference may be shorter than ``symbols`` (e.g. a transmitted preamble
696
+ or the first ``N_ref`` source symbols); estimation uses only the overlapping
697
+ window.
698
+
699
+ Parameters
700
+ ----------
701
+ symbols : array_like
702
+ Equalizer output symbols. Shape: ``(N,)`` or ``(C, N)``.
703
+ ref_symbols : array_like
704
+ Known transmitted symbols. Shape: ``(N_ref,)`` or ``(C, N_ref)``,
705
+ where ``N_ref <= N``. Each channel is matched independently; a
706
+ single-channel ref is broadcast across all output channels.
707
+ num_skip_symbols : int, default 0
708
+ Leading symbols excluded from the rotation estimate (e.g. the
709
+ unconverged equalizer transient). The correction is still applied
710
+ to the full ``symbols``.
711
+
712
+ Returns
713
+ -------
714
+ array_like
715
+ Phase-corrected symbols, same shape and dtype as ``symbols``.
716
+ """
717
+ symbols, xp, _ = dispatch(symbols)
718
+ was_1d = symbols.ndim == 1
719
+ if was_1d:
720
+ symbols = symbols[None, :]
721
+ C, N = symbols.shape
722
+
723
+ ref = xp.asarray(ref_symbols)
724
+ if ref.ndim == 1:
725
+ ref = ref[None, :]
726
+ N_ref = ref.shape[-1]
727
+ if ref.shape[0] == 1 and C > 1:
728
+ ref = xp.broadcast_to(ref, (C, N_ref))
729
+
730
+ if num_skip_symbols >= N_ref:
731
+ raise ValueError(
732
+ f"num_skip_symbols={num_skip_symbols} must be less than the reference "
733
+ f"length N_ref={N_ref}."
734
+ )
735
+
736
+ seg_y = symbols[:, num_skip_symbols:N_ref] # (C, N_est)
737
+ seg_r = ref[:, num_skip_symbols:] # (C, N_est)
738
+ thetas = -xp.angle(xp.sum(seg_y * xp.conj(seg_r), axis=-1)) # (C,) on device
739
+ phasors = xp.exp(1j * thetas).astype(symbols.dtype) # (C,) on device
740
+ out = symbols * phasors[:, None]
741
+
742
+ if logger.isEnabledFor(logging.INFO):
743
+ # Host transfer of thetas is needed only for this per-channel log;
744
+ # the correction itself (phasors) is computed on-device above.
745
+ thetas_deg = np.degrees(to_device(thetas, "cpu"))
746
+ for ch, deg in enumerate(thetas_deg.tolist()):
747
+ logger.info("correct_phase_rotation: ch=%s, theta=%.2f°", ch, deg)
748
+
749
+ if was_1d:
750
+ return out[0]
751
+ return out