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,790 @@
1
+ """Pilot-tone-based polarization demultiplexing."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from collections.abc import Sequence
6
+ from typing import Any, cast
7
+
8
+ import numpy as np
9
+
10
+ from ..backend import ArrayType, dispatch, to_device
11
+ from ..filtering import fir_filter, lowpass_taps
12
+ from ..logger import logger
13
+
14
+
15
+ def apply_interpolated_matrix(
16
+ samples: ArrayType,
17
+ matrix_grid: ArrayType,
18
+ grid_positions: ArrayType,
19
+ ) -> ArrayType:
20
+ r"""Apply a time-varying matrix ``M(n)`` to ``samples``, interpolated from a grid.
21
+
22
+ Output sample ``n`` is ``M(n) · samples[:, n]``, where ``M(n)`` is the exact
23
+ linear blend ``M[g] + (M[g+1]-M[g])·t`` inside grid cell ``g`` (``t`` ramps
24
+ ``0->1`` across the cell). This is the apply half of the dynamic pilot demux,
25
+ but it is a generic operator: given any per-grid-point matrix stack (an inverse
26
+ ``W(n)``, a polar unitary ``Qᴴ(n)``, a butterfly seed, ...) it interpolates and
27
+ applies it as a fixed forward pass.
28
+
29
+ Within each uniform cell the apply collapses to two batched ``(K, C)`` GEMMs -
30
+ ``M[g]·X`` and ``dM·(X·t)`` - instead of materialising a per-sample
31
+ ``(L, K, C)`` matrix and an einsum mat-vec, so the cost is a handful of batched
32
+ kernels rather than an ``O(N)`` Python/launch-bound loop. The ``O(N)`` data
33
+ path runs in ``complex64`` (a well-conditioned mix with no long accumulation),
34
+ so a ``matrix_grid`` built by a double-precision inverse/SVD keeps its
35
+ precision while the bulk traffic moves at half the bandwidth.
36
+
37
+ Parameters
38
+ ----------
39
+ samples : (C, N) array
40
+ Input channels.
41
+ matrix_grid : (G, K, C) array
42
+ Per-grid-point matrices mapping the ``C`` inputs to ``K`` outputs.
43
+ grid_positions : (G,) array
44
+ Sample indices at which ``matrix_grid`` was evaluated. The interior must
45
+ be uniformly spaced (the last point may be pinned to ``N-1``).
46
+
47
+ Returns
48
+ -------
49
+ (K, N) array
50
+ ``M(n) · samples[:, n]``, same dtype as ``samples``.
51
+ """
52
+ samples, xp, _ = dispatch(samples)
53
+ C, N = samples.shape
54
+ M = xp.asarray(matrix_grid, dtype=xp.complex64) # (G, K, C)
55
+ gp = xp.asarray(grid_positions, dtype=xp.float64) # (G,)
56
+ G, K = int(M.shape[0]), int(M.shape[1])
57
+ xc = samples.astype(xp.complex64, copy=False)
58
+ step = int(round(float(gp[1] - gp[0]))) if G > 1 else N
59
+ out = xp.empty((K, N), dtype=xp.complex64)
60
+
61
+ nblk = (N - 1) // step if step > 0 else 0 # full uniform cells over [0, nblk·step)
62
+ bulk = nblk * step
63
+ if nblk > 0:
64
+ Xb = xc[:, :bulk].reshape(C, nblk, step).transpose(1, 0, 2) # (nblk, C, step)
65
+ M0 = M[:nblk] # (nblk, K, C)
66
+ dM = M[1 : nblk + 1] - M0
67
+ t = (xp.arange(step, dtype=xp.float32) / step).astype(xp.complex64)
68
+ # Diagonal scaling on the sample axis commutes through the left matmul
69
+ # (dM @ (Xb·diag(t)) == (dM @ Xb)·diag(t)), so stacking [M0; dM] turns the
70
+ # two batched GEMMs into one and moves the ramp multiply from the
71
+ # input-sized (nblk, C, step) to the output-sized (nblk, K, step) array.
72
+ Y = xp.concatenate([M0, dM], axis=1) @ Xb # (nblk, 2K, step)
73
+ y = Y[:, K:, :]
74
+ y *= t[None, None, :]
75
+ y += Y[:, :K, :]
76
+ out[:, :bulk] = y.transpose(1, 0, 2).reshape(K, bulk)
77
+ if bulk < N: # tail (< step; spans the pinned last cell) - per-sample blend
78
+ nn = xp.arange(bulk, N, dtype=xp.float64)
79
+ lo = xp.clip(xp.searchsorted(gp, nn, side="right") - 1, 0, G - 2)
80
+ frac = ((nn - gp[lo]) / (gp[lo + 1] - gp[lo])).astype(xp.complex64)
81
+ M_full = M[lo] + (M[lo + 1] - M[lo]) * frac[:, None, None] # (L, K, C)
82
+ out[:, bulk:] = xp.einsum("lkc,cl->kl", M_full, xc[:, bulk:])
83
+ return out.astype(samples.dtype, copy=False)
84
+
85
+
86
+ # -----------------------------------------------------------------------------
87
+ # TONE-BASED POLARIZATION DEMULTIPLEXING
88
+ # -----------------------------------------------------------------------------
89
+
90
+ _EXTRACT_CHUNK = 1 << 20 # samples per block in the chunked tone-phasor GEMM
91
+
92
+
93
+ def _tone_phasor_matrix(xw: ArrayType, freqs, sampling_rate: float) -> ArrayType:
94
+ r"""Tone-phasor matrix ``T[i, j] = (1/N) Σ_n xw[i, n]·exp(-j2π f_j n/fs)``.
95
+
96
+ The whole-record accumulation is precision-sensitive (it feeds a matrix
97
+ inverse), but running the O(N) GEMM in complex128 would put the entire data
98
+ path on the slow FP64 units. Instead the record is processed in
99
+ ``_EXTRACT_CHUNK`` blocks: each block's ``(C, chunk) @ (chunk, K)`` product
100
+ runs in the signal's working precision (complex64 for complex64 input) and
101
+ the small ``(C, K)`` per-block partials are accumulated in complex128 - the
102
+ round-off no longer grows with N beyond a block, at complex64 bandwidth.
103
+ The phase ramp is always formed and wrapped in float64 (a float32 ramp
104
+ loses the integer turn count over long records) and only the *wrapped*
105
+ phase drops to float32 for the ``exp``.
106
+ """
107
+ xw, xp, _ = dispatch(xw)
108
+ C, N = xw.shape
109
+ fc = xp.asarray([float(f) for f in freqs], dtype=xp.float64).reshape(-1, 1)
110
+ K = int(fc.shape[0])
111
+ real_dtype = xp.float64 if xw.dtype == xp.complex128 else xp.float32
112
+ two_pi = 2.0 * np.pi
113
+ T = xp.zeros((C, K), dtype=xp.complex128)
114
+ for start in range(0, N, _EXTRACT_CHUNK):
115
+ stop = min(start + _EXTRACT_CHUNK, N)
116
+ nn = xp.arange(start, stop, dtype=xp.float64)
117
+ ph = fc * (two_pi / sampling_rate) * nn[None, :] # (K, chunk) float64
118
+ ph -= xp.round(ph / two_pi) * two_pi
119
+ basis = xp.exp(-1j * ph.astype(real_dtype)) # (K, chunk) working dtype
120
+ T += xw[:, start:stop] @ basis.T # (C, K) partial, accumulated in c128
121
+ return T / N
122
+
123
+
124
+ def _refine_tone_frequencies(
125
+ xw: ArrayType,
126
+ T: ArrayType,
127
+ freqs,
128
+ sampling_rate: float,
129
+ search_band: float,
130
+ ) -> list[float]:
131
+ """Sub-bin refine each tone on the receive channel where it is strongest.
132
+
133
+ Two stages, replacing the per-tone ``find_bias_tone`` calls (each of which
134
+ ran its own full-record, power-of-two zero-padded FFT):
135
+
136
+ 1. **Coarse**: one batched FFT (working precision) of only the *unique*
137
+ strongest channels + the device-side log-parabolic fit - good to a
138
+ fraction of a bin, but ill-conditioned for a tone sitting exactly on a
139
+ bin (its neighbours are sinc nulls).
140
+ 2. **Fine**: two-segment phase slope - the exact tone phasor of each
141
+ half-record at the coarse frequency (chunked GEMMs); the phase advance
142
+ between the segment centroids gives the residual offset
143
+ ``δ = ∠(z_b·conj(z_a))·fs/(2π·Δc)``, unambiguous for ``|δ| < 1`` bin
144
+ (which stage 1 guarantees) and accurate to well below either parabolic
145
+ fit. The demux mix-down is Hz-sensitive - a tone-frequency error
146
+ leaves a slow residual rotation on the demuxed output - hence this
147
+ stage.
148
+ """
149
+ from ..frequency import _refine_tones_from_spectrum
150
+
151
+ xw, xp, _ = dispatch(xw)
152
+ N = int(xw.shape[-1])
153
+ K = len(list(freqs))
154
+ best_ch = to_device(xp.argmax(xp.abs(T), axis=0), "cpu") # (K,) host ints
155
+ uniq, inv = np.unique(np.asarray(best_ch), return_inverse=True)
156
+ xr = xw[xp.asarray(uniq)] # (U, N) unique strongest channels only
157
+ X = xp.fft.fft(xr, axis=-1)
158
+ f_coarse = _refine_tones_from_spectrum(
159
+ X, sampling_rate, freqs, search_band, rows=inv
160
+ ) # (K,) float64 host
161
+
162
+ # Two-segment phase slope. z_b is computed with a local time origin, so
163
+ # its known carrier phase 2π·f̂·N2/fs is undone first - evaluated in
164
+ # *turns* on host float64 (the fractional part keeps ~1e-9 rad precision
165
+ # for any realistic record length).
166
+ N2 = N // 2
167
+ Za = _tone_phasor_matrix(xr[:, :N2], f_coarse, sampling_rate) # (U, K)
168
+ Zb = _tone_phasor_matrix(xr[:, N2:], f_coarse, sampling_rate) # (U, K)
169
+ rows_dev = xp.asarray(inv)
170
+ cols = xp.arange(K)
171
+ turns = f_coarse * N2 / sampling_rate
172
+ corr = xp.asarray(np.exp(-2j * np.pi * (turns - np.round(turns)))) # (K,)
173
+ dphi = to_device(
174
+ xp.angle(Zb[rows_dev, cols] * corr * xp.conj(Za[rows_dev, cols])), "cpu"
175
+ )
176
+ dc = N2 + 0.5 * (N % 2) # exact centroid spacing of the two segments
177
+ delta_f = dphi * sampling_rate / (2.0 * np.pi * dc)
178
+ return [float(f) for f in f_coarse + delta_f]
179
+
180
+
181
+ def _jones_at_grid_points(
182
+ xw: ArrayType,
183
+ h: np.ndarray,
184
+ freqs,
185
+ grid_np: np.ndarray,
186
+ sampling_rate: float,
187
+ ) -> ArrayType:
188
+ r"""Tone-tracked Jones estimate evaluated **only** at the grid points.
189
+
190
+ The mix-down + centred-FIR tracker output at sample ``g`` factors as
191
+
192
+ T[c, k, g] = e^{-j2π f_k g/fs} · Σ_q h̃_k[q] · r_c[g - lead + q],
193
+ h̃_k[q] = h[L-1-q] · e^{-j2π f_k (q - lead)/fs}, lead = ⌈(L-1)/2⌉,
194
+
195
+ i.e. a bandpass filter at ``+f_k`` - so instead of materialising the
196
+ (C, K, N) mixed array and FFT-filtering every sample only to keep one in
197
+ ``grid_step``, this gathers the (chunk, L) sample windows around the grid
198
+ points and hits all K modulated tap vectors with one GEMM per chunk. No
199
+ full-record temporaries are created. Zero-padding at the record edges
200
+ reproduces ``mode='same'`` semantics exactly, so the ``num_taps//2``
201
+ edge-guard contract is unchanged.
202
+
203
+ Parameters
204
+ ----------
205
+ xw : (C, N) array
206
+ Received samples in working precision (complex64/complex128).
207
+ h : (L,) np.ndarray
208
+ Tracking low-pass FIR taps.
209
+ freqs : sequence of float
210
+ The K tracked tone frequencies in Hz.
211
+ grid_np : (G,) np.ndarray of int
212
+ Sample indices of the inversion grid (host).
213
+ sampling_rate : float
214
+ Sampling rate in Hz.
215
+
216
+ Returns
217
+ -------
218
+ (G, C, K) array
219
+ Jones stack in complex128, ready for the batched Gram inverse.
220
+ """
221
+ xw, xp, _ = dispatch(xw)
222
+ C, N = xw.shape
223
+ L = int(len(h))
224
+ G = int(len(grid_np))
225
+ lead = (L - 1) - (L - 1) // 2
226
+ real_dtype = xp.float64 if xw.dtype == xp.complex128 else xp.float32
227
+ two_pi = 2.0 * np.pi
228
+ f_dev = xp.asarray([float(f) for f in freqs], dtype=xp.float64) # (K,)
229
+ K = int(f_dev.shape[0])
230
+
231
+ # Modulated (bandpass) taps h̃ (L, K) and per-grid-point de-rotation (G, K):
232
+ # phases formed and wrapped in float64, exp in working precision.
233
+ q = xp.arange(L, dtype=xp.float64)
234
+ ph_m = (two_pi / sampling_rate) * (q[:, None] - lead) * f_dev[None, :] # (L, K)
235
+ ph_m -= xp.round(ph_m / two_pi) * two_pi
236
+ h_rev = xp.asarray(np.ascontiguousarray(h[::-1]), dtype=real_dtype) # h[L-1-q]
237
+ Ht = h_rev[:, None] * xp.exp(-1j * ph_m.astype(real_dtype)) # (L, K)
238
+
239
+ gp = xp.asarray(grid_np, dtype=xp.float64)
240
+ ph_g = (two_pi / sampling_rate) * gp[:, None] * f_dev[None, :] # (G, K)
241
+ ph_g -= xp.round(ph_g / two_pi) * two_pi
242
+ rot = xp.exp(-1j * ph_g.astype(real_dtype)) # (G, K)
243
+
244
+ # Interior windows lie fully inside [0, N); the few edge windows get a
245
+ # masked (zero-padded) gather. Chunk the gather so the (chunk, L) window
246
+ # and index arrays stay ~128 MB regardless of G·L.
247
+ starts_np = grid_np.astype(np.int64) - lead # (G,) host, sorted
248
+ i0 = int(np.searchsorted(starts_np, 0, side="left"))
249
+ i1 = int(np.searchsorted(starts_np, N - L, side="right"))
250
+ idx_dtype = xp.int64 if (N + L) >= (1 << 31) else xp.int32
251
+ itemsize = xw.dtype.itemsize
252
+ chunk = max(1, int((128 << 20) // (L * (itemsize + 4))))
253
+ off = xp.arange(L, dtype=idx_dtype)
254
+ Tg_work = xp.empty((C, G, K), dtype=xw.dtype)
255
+
256
+ def _fill(a: int, b: int, masked: bool) -> None:
257
+ for s0 in range(a, b, chunk):
258
+ s1 = min(s0 + chunk, b)
259
+ st = xp.asarray(starts_np[s0:s1], dtype=idx_dtype)
260
+ idx = st[:, None] + off[None, :] # (chunk, L)
261
+ if masked:
262
+ valid = (idx >= 0) & (idx < N)
263
+ idx = xp.clip(idx, 0, N - 1)
264
+ for c in range(C):
265
+ wnd = xw[c][idx] # (chunk, L) gather
266
+ if masked:
267
+ wnd *= valid
268
+ Tg_work[c, s0:s1] = wnd @ Ht # (chunk, K)
269
+
270
+ _fill(0, i0, True) # leading edge (zero-padded)
271
+ _fill(i0, i1, False) # interior - plain gather
272
+ _fill(i1, G, True) # trailing edge
273
+
274
+ Tg_work *= rot[None, :, :]
275
+ return xp.moveaxis(Tg_work, 0, 1).astype(xp.complex128) # (G, C, K)
276
+
277
+
278
+ def demultiplex_polarization_tones_static(
279
+ samples: ArrayType,
280
+ sampling_rate: float,
281
+ tone_frequencies: Sequence[float],
282
+ *,
283
+ refine_tones: bool = True,
284
+ search_band: float | None = None,
285
+ normalize: bool = True,
286
+ return_matrix: bool = False,
287
+ ) -> ArrayType | tuple[ArrayType, ArrayType]:
288
+ r"""
289
+ One-shot polarization demux from distinct per-stream CW pilot tones.
290
+
291
+ Undoes a **frequency-flat** polarization / spatial mixing by inverting the
292
+ channel's Jones matrix, which is read directly off pilot tones placed at a
293
+ *distinct* frequency on each transmitted stream (see ``add_pilot_tone`` with
294
+ a per-channel ``frequency`` list).
295
+
296
+ **Principle.** With tone ``f_j`` carried *only* by transmitted stream ``j``
297
+ and a frequency-flat mixing ``r = J s`` (e.g.
298
+ ``apply_polarization_mixing``), the complex amplitude of tone ``f_j`` in
299
+ received channel ``i`` is ``T[i, j] = J[i, j] · α_j`` (``α_j`` = the TX tone
300
+ amplitude of stream ``j``). Hence the measured tone-phasor matrix factors as
301
+ ``T = J · diag(α)``, and ``W = pinv(T)`` unmixes::
302
+
303
+ W r = diag(1/α) · J^{-1} J s = diag(1/α) · s,
304
+
305
+ recovering each stream up to a trivial per-stream complex scale ``1/α_j``
306
+ (removed by ``normalize`` and/or downstream CPR). Because each tone uniquely
307
+ labels its stream, output row ``j`` always corresponds to ``tone_frequencies[j]``
308
+ - there is **no polarization-permutation ambiguity** (unlike blind CMA; cf.
309
+ ``resolve_polarization_permutation``).
310
+
311
+ **Speed.** Tones added with ``add_pilot_tone`` are grid-quantized
312
+ (buffer-periodic), so a single DFT bin is the exact, mutually-orthogonal,
313
+ maximum-likelihood estimator of each tone phasor. The whole operation is two
314
+ small GEMMs (extract ``T`` via a ``(C,N)·(N,K)`` product, apply ``W`` via a
315
+ ``(K,C)·(C,N)`` product) plus one ``KxK`` inverse - no iteration, no
316
+ convergence, far below the cost of a full FFT.
317
+
318
+ **Scope.** Frequency-flat **and time-invariant** SOP only - the whole-record
319
+ DFT bin estimates a *single* Jones matrix. If the state of polarization
320
+ drifts appreciably over the capture (long records and/or low baud rate, so
321
+ the wall-clock duration exceeds the SOP coherence time), the averaged tone
322
+ phasor is biased and attenuated and the one-shot inverse leaves residual,
323
+ time-growing crosstalk; use ``demultiplex_polarization_tones_dynamic``
324
+ instead. For PMD / DGD use the returned unmixer (``return_matrix=True``) to
325
+ seed a butterfly equalizer (``cma`` / ``block_lms``).
326
+
327
+ Parameters
328
+ ----------
329
+ samples : array_like
330
+ Received MIMO samples. Shape ``(C, N)`` - time on the last axis.
331
+ sampling_rate : float
332
+ Sampling rate f_s in Hz.
333
+ tone_frequencies : sequence of float
334
+ The ``K`` distinct per-stream tone frequencies in Hz (as added at the
335
+ TX, in transmitted-stream order). Require ``K <= C``. Output row ``j``
336
+ corresponds to ``tone_frequencies[j]``.
337
+ refine_tones : bool, default True
338
+ If ``True``, sub-bin-refine each tone centre with ``find_bias_tone``
339
+ (on the receive channel where that tone is strongest) before extraction,
340
+ absorbing a residual carrier frequency offset that has dragged the tone
341
+ off its nominal bin. If ``False``, extract exactly at
342
+ ``tone_frequencies``.
343
+ search_band : float, optional
344
+ Half-width in Hz of the per-tone peak search when ``refine_tones=True``.
345
+ Defaults to ``4 * f_s / N`` (a few FFT bins). Widen it if the carrier
346
+ offset can exceed that, but keep it inside the tone-to-data guard so the
347
+ data band never wins the argmax.
348
+ normalize : bool, default True
349
+ If ``True``, rescale each demuxed row so its mean power equals the mean
350
+ per-channel input power, removing the arbitrary ``1/α_j`` per-stream
351
+ scale and preserving the library power invariant.
352
+ return_matrix : bool, default False
353
+ If ``True``, also return the ``(K, C)`` unmixing matrix ``W``.
354
+
355
+ Returns
356
+ -------
357
+ demuxed : array_like
358
+ Demultiplexed streams. Shape ``(K, N)`` - one row per recovered
359
+ *stream* (``K`` = number of tones), **not** per receive channel: the
360
+ ``C`` input channels are mapped down to the ``K`` transmitted streams.
361
+ In the usual square dual-pol case ``K == C == 2`` the two coincide.
362
+ Same complex dtype and backend as the input. Row ``j`` is the stream
363
+ that carried ``tone_frequencies[j]``.
364
+ W : array_like, optional
365
+ Returned only if ``return_matrix=True``: the ``(K, C)`` unmixing matrix
366
+ (``complex128``), i.e. the estimated inverse Jones matrix up to per-stream
367
+ scaling. Suitable as a seed for a butterfly equalizer.
368
+
369
+ Raises
370
+ ------
371
+ ValueError
372
+ If ``samples`` is not 2-D, ``tone_frequencies`` is empty, ``K > C``, or
373
+ any tone frequency lies outside ``(-f_s/2, f_s/2)``.
374
+
375
+ See Also
376
+ --------
377
+ add_pilot_tone : Add the per-stream tones at the transmitter.
378
+ demultiplex_polarization_tones_dynamic : Time-varying (drifting-SOP) demux.
379
+ """
380
+ samples, xp, _ = dispatch(samples)
381
+ if samples.ndim != 2:
382
+ raise ValueError(
383
+ "demultiplex_polarization_tones_static requires a 2-D (C, N) MIMO "
384
+ f"input; got ndim={samples.ndim}."
385
+ )
386
+ C, N = samples.shape
387
+
388
+ f_tones = [float(f) for f in tone_frequencies]
389
+ K = len(f_tones)
390
+ if K == 0:
391
+ raise ValueError("tone_frequencies must contain at least one frequency.")
392
+ if K > C:
393
+ raise ValueError(
394
+ f"got K={K} tones but only C={C} receive channels; need K <= C to "
395
+ "unmix (one tone per transmitted stream)."
396
+ )
397
+ nyq = sampling_rate / 2.0
398
+ for f in f_tones:
399
+ if not (-nyq < f < nyq):
400
+ raise ValueError(
401
+ f"tone_frequencies entry {f} must lie in (-fs/2, fs/2) = "
402
+ f"(±{nyq:.3g}) Hz."
403
+ )
404
+
405
+ # The KxK inverse is precision-sensitive (CLAUDE.md) and stays in
406
+ # complex128, but every O(N) pass runs in the signal's working precision:
407
+ # the tone phasors are accumulated block-wise with complex128 partials
408
+ # (_tone_phasor_matrix), and the unmix is a well-conditioned per-sample
409
+ # mat-vec with no long accumulation, so complex64 halves the bulk traffic
410
+ # without touching the precision that matters.
411
+ xw = (
412
+ samples
413
+ if samples.dtype == xp.complex128
414
+ else samples.astype(xp.complex64, copy=False)
415
+ )
416
+
417
+ df = sampling_rate / N
418
+ T = _tone_phasor_matrix(xw, f_tones, sampling_rate) # (C, K) complex128
419
+
420
+ if refine_tones:
421
+ if search_band is None:
422
+ search_band = 4.0 * df
423
+ # Sub-bin refine each tone on the channel where it is strongest (one
424
+ # batched FFT of the unique channels; single host transfer), re-extract.
425
+ f_used = _refine_tone_frequencies(xw, T, f_tones, sampling_rate, search_band)
426
+ T = _tone_phasor_matrix(xw, f_used, sampling_rate)
427
+ else:
428
+ f_used = f_tones
429
+
430
+ # Unmix. pinv covers the over-determined C > K case and equals the inverse
431
+ # when square; kept in complex128 (inversion is precision-sensitive).
432
+ W = xp.linalg.pinv(T) # (K, C)
433
+ demuxed = W.astype(xw.dtype) @ xw # (K, N) working precision
434
+
435
+ if normalize:
436
+ # Single-pass BLAS reduction - no full-record |x|² temporary.
437
+ p_in = xp.vdot(samples, samples).real / samples.size
438
+ p_out = xp.mean(xp.abs(demuxed) ** 2, axis=-1, keepdims=True) # (K, 1)
439
+ scale = xp.sqrt(p_in / xp.where(p_out > 0, p_out, 1.0))
440
+ demuxed = demuxed * scale
441
+
442
+ demuxed = demuxed.astype(samples.dtype, copy=False)
443
+
444
+ logger.info(
445
+ "demultiplex_polarization_tones: f_tones=%s Hz, refine=%s [C=%d, K=%d, N=%d]",
446
+ [f"{f:.3g}" for f in f_used],
447
+ refine_tones,
448
+ C,
449
+ K,
450
+ N,
451
+ )
452
+
453
+ if return_matrix:
454
+ return demuxed, W
455
+ return demuxed
456
+
457
+
458
+ def demultiplex_polarization_tones_dynamic(
459
+ samples: ArrayType,
460
+ sampling_rate: float,
461
+ tone_frequencies: Sequence[float],
462
+ *,
463
+ track_bandwidth: float,
464
+ num_taps: int | None = None,
465
+ grid_step: int | None = None,
466
+ refine_tones: bool = True,
467
+ search_band: float | None = None,
468
+ normalize: bool = True,
469
+ trim_edges: bool = False,
470
+ return_matrix: bool = False,
471
+ apply: bool = True,
472
+ ) -> ArrayType | tuple[Any, ...]:
473
+ r"""
474
+ Time-varying polarization demux from distinct per-stream CW pilot tones.
475
+
476
+ Drifting-SOP counterpart of ``demultiplex_polarization_tones_static``. Where
477
+ the static routine reads a *single* Jones matrix from a whole-record DFT bin,
478
+ this one **tracks** a slowly rotating frequency-flat mixing ``r(n) = J(n) s(n)``
479
+ by following each pilot tone continuously in time.
480
+
481
+ **Principle.** Tone ``f_j`` is a CW carried *only* by transmitted stream
482
+ ``j`` (amplitude ``α_j``). Mixing receive channel ``i`` down by ``f_j`` and
483
+ low-pass filtering isolates that tone's slowly-varying phasor::
484
+
485
+ z_ij(n) = LPF{ r_i(n) · exp(-j2π f_j n / f_s) } ≈ J_ij(n) · α_j,
486
+
487
+ because every other tone ``f_k`` (k ≠ j) lands at ``f_k - f_j`` and the data
488
+ band is pushed away from DC, so the LPF rejects them. Running this for all
489
+ ``K`` tones (one mix-down per *distinct* tone frequency) yields a continuous
490
+ estimate of the whole Jones matrix ``T(n) = J(n) diag(α)``, shape ``(C, K, N)``.
491
+ Inverting ``T`` on a decimated time grid and interpolating back to full rate
492
+ gives a per-sample unmixer ``W(n) = pinv(T(n))`` with::
493
+
494
+ W(n) r(n) = diag(1/α) · s(n),
495
+
496
+ recovering each stream up to the same trivial per-stream scale ``1/α_j`` as
497
+ the static routine. As with the static version each tone uniquely labels its
498
+ stream, so there is **no polarization-permutation ambiguity**: output row
499
+ ``j`` corresponds to ``tone_frequencies[j]``.
500
+
501
+ **Tracking-bandwidth trade-off.** ``track_bandwidth`` (the LPF cut-off) is
502
+ the single design knob and is bounded on both sides:
503
+
504
+ * It must be **≥ the SOP rotation rate**, or ``W(n)`` lags the true ``J(n)``
505
+ and residual crosstalk remains (lag bias).
506
+ * It must be **≤ the guard** to the nearest other tone and to the data band,
507
+ or those leak into ``z_ij`` and corrupt the estimate. The hard ceiling is
508
+ roughly half the smallest tone spacing.
509
+
510
+ If the SOP rotates faster than the available tone spacing permits to track,
511
+ the tones are simply spaced too closely for that drift - a real feasibility
512
+ limit; a warning is logged when ``2·track_bandwidth`` (plus the FIR
513
+ transition) encroaches on the nearest tone spacing.
514
+
515
+ Parameters
516
+ ----------
517
+ samples : array_like
518
+ Received MIMO samples. Shape ``(C, N)`` - time on the last axis.
519
+ sampling_rate : float
520
+ Sampling rate f_s in Hz.
521
+ tone_frequencies : sequence of float
522
+ The ``K`` distinct per-stream tone frequencies in Hz (as added at the
523
+ TX, in transmitted-stream order). Require ``K <= C``. Output row ``j``
524
+ corresponds to ``tone_frequencies[j]``.
525
+ track_bandwidth : float
526
+ One-sided LPF cut-off in Hz - the polarization-tracking bandwidth. Set
527
+ it a few times above the expected SOP rotation rate but well below the
528
+ smallest tone spacing (see the trade-off above).
529
+ num_taps : int, optional
530
+ Length of the tracking low-pass FIR. Defaults to ``~3.3·f_s/track_bandwidth``
531
+ (the Hamming transition width that resolves ``track_bandwidth``), forced
532
+ odd and clipped below ``N``. Increase for sharper neighbour-tone
533
+ rejection at the cost of longer edge transients.
534
+ grid_step : int, optional
535
+ Decimation (in samples) of the grid on which ``T(n)`` is inverted.
536
+ Defaults to ``max(1, floor(f_s / (4·track_bandwidth)))`` - i.e. oversample
537
+ the tracked process ~4x. ``W`` is linearly interpolated between grid
538
+ points, so a finer grid costs more inverses but tracks marginally better.
539
+ refine_tones : bool, default True
540
+ If ``True``, sub-bin-refine each tone centre with ``find_bias_tone`` (on
541
+ the receive channel where it is strongest) before mixing down, absorbing
542
+ a residual carrier frequency offset.
543
+ search_band : float, optional
544
+ Half-width in Hz of the per-tone peak search when ``refine_tones=True``.
545
+ Defaults to ``4 · f_s / N``.
546
+ normalize : bool, default True
547
+ If ``True``, rescale each demuxed row so its mean power equals the mean
548
+ per-channel input power (removes the ``1/α_j`` scale; preserves the
549
+ library power invariant). When ``trim_edges=True`` the power is measured
550
+ over the retained interior only.
551
+ trim_edges : bool, default False
552
+ The tracking FIR is applied with centred (``'same'``) convolution, so the
553
+ Jones estimate - and hence ``W(n)`` - is unreliable within ``num_taps//2``
554
+ samples of each record end (the convolution averages in zero-padding
555
+ there). The **data is never filtered**, so timing is unaffected, but
556
+ those edge samples carry residual crosstalk. If ``True``, drop them:
557
+ ``demuxed`` is returned as the reliable interior ``(K, N - 2·g)`` with
558
+ ``g = num_taps//2``, together with a ``valid`` slice giving the retained
559
+ sample range in **original** coordinates (so full-length references align
560
+ as ``ref[..., valid]``).
561
+ return_matrix : bool, default False
562
+ If ``True``, also return the decimated unmixer stack ``W_grid`` and the
563
+ sample positions ``grid_positions`` it was evaluated at (suitable for
564
+ seeding a time-varying butterfly equalizer). ``W_grid`` / ``grid_positions``
565
+ always span the **full** record, even when ``trim_edges=True``.
566
+ apply : bool, default True
567
+ If ``True`` (default), interpolate ``W(n)`` to full rate and apply it,
568
+ returning the demuxed signal as documented below. If ``False``,
569
+ **matrix-only mode**: skip the ``O(N)`` interpolate-and-apply entirely and
570
+ return just ``(W_grid, grid_positions)`` (``return_matrix`` is implied).
571
+ Use this when only the unmixer stack is needed - e.g. to make a
572
+ PDL/unitarity decision and then apply a *different* factor (a polar unitary
573
+ ``Qᴴ(n)``) without paying for a demux that would be discarded. ``normalize``
574
+ and ``trim_edges`` act on the applied signal, so they have **no effect**
575
+ when ``apply=False``.
576
+
577
+ Returns
578
+ -------
579
+ demuxed : array_like
580
+ Demultiplexed streams. Same complex dtype and backend as the input; row
581
+ ``j`` carried ``tone_frequencies[j]``. Shape ``(K, N)``, or
582
+ ``(K, N - 2·(num_taps//2))`` when ``trim_edges=True``. **Omitted** when
583
+ ``apply=False`` (the return is then ``(W_grid, grid_positions)``).
584
+ valid : slice, optional
585
+ Returned only if ``trim_edges=True``: the ``slice(g, N - g)`` of original
586
+ sample indices retained in ``demuxed`` (``g = num_taps//2``). Always
587
+ precedes ``W_grid`` in the output tuple.
588
+ W_grid : array_like, optional
589
+ Returned only if ``return_matrix=True``: the ``(G, K, C)`` stack of
590
+ per-grid-point unmixing matrices (``complex128``).
591
+ grid_positions : array_like, optional
592
+ Returned only if ``return_matrix=True``: the ``(G,)`` sample indices
593
+ (``float64``) at which ``W_grid`` was evaluated.
594
+
595
+ Raises
596
+ ------
597
+ ValueError
598
+ If ``samples`` is not 2-D, ``tone_frequencies`` is empty, ``K > C``,
599
+ ``track_bandwidth`` is not positive, or any tone frequency lies outside
600
+ ``(-f_s/2, f_s/2)``.
601
+
602
+ See Also
603
+ --------
604
+ demultiplex_polarization_tones_static : One-shot static-SOP demux (faster).
605
+ add_pilot_tone : Add the per-stream tones at the transmitter.
606
+ """
607
+ samples, xp, _ = dispatch(samples)
608
+ if samples.ndim != 2:
609
+ raise ValueError(
610
+ "demultiplex_polarization_tones_dynamic requires a 2-D (C, N) MIMO "
611
+ f"input; got ndim={samples.ndim}."
612
+ )
613
+ C, N = samples.shape
614
+
615
+ f_tones = [float(f) for f in tone_frequencies]
616
+ K = len(f_tones)
617
+ if K == 0:
618
+ raise ValueError("tone_frequencies must contain at least one frequency.")
619
+ if K > C:
620
+ raise ValueError(
621
+ f"got K={K} tones but only C={C} receive channels; need K <= C to "
622
+ "unmix (one tone per transmitted stream)."
623
+ )
624
+ if not (track_bandwidth > 0):
625
+ raise ValueError(f"track_bandwidth must be positive; got {track_bandwidth}.")
626
+ nyq = sampling_rate / 2.0
627
+ for f in f_tones:
628
+ if not (-nyq < f < nyq):
629
+ raise ValueError(
630
+ f"tone_frequencies entry {f} must lie in (-fs/2, fs/2) = "
631
+ f"(±{nyq:.3g}) Hz."
632
+ )
633
+
634
+ df = sampling_rate / N
635
+ # Working precision for all O(N) traffic (complex64 unless the caller
636
+ # supplied complex128); the batched inverse below stays complex128.
637
+ xw = (
638
+ samples
639
+ if samples.dtype == xp.complex128
640
+ else samples.astype(xp.complex64, copy=False)
641
+ )
642
+
643
+ # --- Tracking low-pass design + feasibility check ------------------------
644
+ if num_taps is None:
645
+ # Hamming transition width ≈ 3.3·fs/num_taps; size it to resolve the
646
+ # requested tracking bandwidth. Force odd; keep it shorter than N.
647
+ num_taps = int(round(3.3 * sampling_rate / track_bandwidth))
648
+ num_taps += 1 - (num_taps % 2) # nearest odd >= value
649
+ num_taps = max(num_taps, 3)
650
+ num_taps = min(int(num_taps), (N // 2) * 2 - 1)
651
+ h = lowpass_taps(sampling_rate, num_taps, track_bandwidth)
652
+
653
+ # Edge guard: 'same' convolution corrupts num_taps//2 samples at each end.
654
+ # num_taps is clipped < N above, so the retained interior is always non-empty.
655
+ guard = num_taps // 2 if trim_edges else 0
656
+
657
+ if K > 1:
658
+ sorted_f = sorted(f_tones)
659
+ d_min = min(b - a for a, b in zip(sorted_f, sorted_f[1:]))
660
+ transition = 3.3 * sampling_rate / num_taps
661
+ if 2.0 * track_bandwidth + transition >= d_min:
662
+ logger.warning(
663
+ "demultiplex_polarization_tones_dynamic: tracking bandwidth "
664
+ "(%.3g Hz, FIR transition %.3g Hz) approaches the nearest tone "
665
+ "spacing %.3g Hz - neighbouring tones may leak into the Jones "
666
+ "estimate. Reduce track_bandwidth or widen the tone spacing.",
667
+ track_bandwidth,
668
+ transition,
669
+ d_min,
670
+ )
671
+
672
+ # --- Optional sub-bin tone refinement (static one-shot bin picks, per
673
+ # tone, the receive channel where it is strongest) --------------------
674
+ if refine_tones:
675
+ if search_band is None:
676
+ search_band = 4.0 * df
677
+ T0 = _tone_phasor_matrix(xw, f_tones, sampling_rate) # (C, K)
678
+ f_used = _refine_tone_frequencies(xw, T0, f_tones, sampling_rate, search_band)
679
+ else:
680
+ f_used = f_tones
681
+
682
+ # --- Decimated inversion grid --------------------------------------------
683
+ if grid_step is None:
684
+ grid_step = max(1, int(sampling_rate / (4.0 * track_bandwidth)))
685
+ grid_step = int(min(max(grid_step, 1), N))
686
+ grid_np = np.arange(0, N, grid_step, dtype=np.int64)
687
+ if int(grid_np[-1]) != N - 1:
688
+ grid_np = np.concatenate([grid_np, [N - 1]]) # pin the last sample
689
+ G = int(grid_np.shape[0])
690
+ grid_positions = xp.asarray(grid_np, dtype=xp.float64)
691
+
692
+ # --- Track the Jones matrix at the grid points ---------------------------
693
+ # T[i, j, g] = LPF{ r_i(n) · exp(-j2π f_j n/fs) }(g) ≈ J_ij(g)·α_j, with a
694
+ # centred ('same') linear-phase FIR, so the LPF group delay is compensated
695
+ # and the estimate aligns in time with the input. Only the G grid points
696
+ # feed the batched inverse, so the tracker is evaluated there directly
697
+ # (windowed gather + GEMM in _jones_at_grid_points) instead of filtering
698
+ # all N samples and discarding grid_step-1 of every grid_step outputs.
699
+ # The batched inverse follows the RLS-style double-precision convention,
700
+ # hence the complex128 (G, C, K) stack.
701
+ if G * num_taps <= 24 * K * N:
702
+ Tg = _jones_at_grid_points(xw, h, f_used, grid_np, sampling_rate)
703
+ else:
704
+ # Dense grid (grid_step ≪ num_taps): per-grid-point windows would touch
705
+ # far more elements than the record itself, so keep the full-rate
706
+ # batched-FIR formulation. The phase ramp 2π·f·n/fs reaches ~1e7 rad
707
+ # over a long capture, so it MUST be formed and wrapped in float64 -
708
+ # float32 loses the integer turn count and corrupts every tone phasor.
709
+ # The wrapped phase, the exp, the mix-down, and the averaging LPF are
710
+ # well-conditioned, so they run in working precision, with FFT-FIR
711
+ # round-off (~√N_fft·ε ≈ 5e-5) far below the demux crosstalk floor.
712
+ two_pi = 2.0 * np.pi
713
+ real_dtype = xp.float64 if xw.dtype == xp.complex128 else xp.float32
714
+ f_arr = xp.asarray([float(fj) for fj in f_used], dtype=xp.float64) # (K,)
715
+ n = xp.arange(N, dtype=xp.float64)
716
+ ph = two_pi * f_arr[:, None] * n[None, :] / sampling_rate # (K, N) float64
717
+ ph -= xp.round(ph / two_pi) * two_pi # wrap in float64 (essential)
718
+ carrier = xp.exp(-1j * ph.astype(real_dtype)) # (K, N) working precision
719
+ mixed = xw[:, None, :] * carrier[None, :, :] # (C, K, N)
720
+ # One batched linear-phase FIR over (C·K) rows instead of K calls.
721
+ T_t = cast(ArrayType, fir_filter(mixed.reshape(C * K, N), h, axis=-1)).reshape(
722
+ C, K, N
723
+ )
724
+ idx = xp.asarray(grid_np)
725
+ Tg = xp.moveaxis(T_t[:, :, idx], 2, 0).astype(xp.complex128) # (G, C, K)
726
+ Th = xp.conj(xp.swapaxes(Tg, -1, -2)) # (G, K, C)
727
+ gram = Th @ Tg # (G, K, K)
728
+ # Tikhonov regularisation keeps the batched inverse well-conditioned when the
729
+ # instantaneous SOP nearly aligns two streams (gram -> singular).
730
+ diag_mean = xp.real(xp.trace(gram, axis1=-2, axis2=-1)) / K # (G,)
731
+ eye = xp.eye(K, dtype=xp.complex128)
732
+ ridge = (1e-9 * diag_mean)[:, None, None] * eye[None, :, :]
733
+ Wg = xp.linalg.inv(gram + ridge) @ Th # (G, K, C) - batched, CPU+GPU
734
+
735
+ if not apply:
736
+ # Matrix-only mode: skip the O(N) interpolate-and-apply. normalize and
737
+ # trim_edges act on the applied signal, so they are no-ops here.
738
+ logger.info(
739
+ "demultiplex_polarization_tones_dynamic (matrix-only): f_tones=%s Hz, "
740
+ "refine=%s, track_bw=%.3g Hz, taps=%d, grid_step=%d, G=%d "
741
+ "[C=%d, K=%d, N=%d]",
742
+ [f"{f:.3g}" for f in f_used],
743
+ refine_tones,
744
+ track_bandwidth,
745
+ num_taps,
746
+ grid_step,
747
+ G,
748
+ C,
749
+ K,
750
+ N,
751
+ )
752
+ return Wg, grid_positions
753
+
754
+ # Interpolate W(n) to full rate and apply it (block-vectorised GEMMs in the
755
+ # shared helper) as a fixed forward pass.
756
+ demuxed = apply_interpolated_matrix(samples, Wg, grid_positions) # (K, N)
757
+
758
+ # Drop the FIR edge transient (the data is untouched; only W is unreliable
759
+ # there). ``valid`` reports the retained range in original coordinates.
760
+ valid = slice(guard, N - guard)
761
+ demuxed = demuxed[:, valid]
762
+
763
+ if normalize:
764
+ p_in = xp.mean(xp.abs(samples[:, valid]) ** 2)
765
+ p_out = xp.mean(xp.abs(demuxed) ** 2, axis=-1, keepdims=True) # (K, 1)
766
+ scale = xp.sqrt(p_in / xp.where(p_out > 0, p_out, 1.0))
767
+ demuxed = demuxed * scale
768
+
769
+ demuxed = demuxed.astype(samples.dtype)
770
+
771
+ logger.info(
772
+ "demultiplex_polarization_tones_dynamic: f_tones=%s Hz, refine=%s, "
773
+ "track_bw=%.3g Hz, taps=%d, grid_step=%d, G=%d [C=%d, K=%d, N=%d]",
774
+ [f"{f:.3g}" for f in f_used],
775
+ refine_tones,
776
+ track_bandwidth,
777
+ num_taps,
778
+ grid_step,
779
+ G,
780
+ C,
781
+ K,
782
+ N,
783
+ )
784
+
785
+ out: tuple[Any, ...] = (demuxed,)
786
+ if trim_edges:
787
+ out = out + (valid,)
788
+ if return_matrix:
789
+ out = out + (Wg, grid_positions)
790
+ return out[0] if len(out) == 1 else out