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,1720 @@
1
+ """JAX (lax.scan) sequential and block equalizer kernels."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import functools
6
+ from typing import Any
7
+
8
+ from ..backend import _get_jax
9
+
10
+ # -----------------------------------------------------------------------------
11
+ # JAX KERNELS - ADAPTIVE EQUALIZER SCANS
12
+ # -----------------------------------------------------------------------------
13
+ #
14
+ # Each factory JIT-compiles a jax.lax.scan kernel on first call and caches
15
+ # it in _JITTED_EQ. Static closure variables are baked into the compiled XLA
16
+ # program - changing any of them produces a cache miss and triggers a new
17
+ # trace+compile (not a retrace of an existing graph).
18
+ #
19
+ # Static variables (shared by all three kernels):
20
+ # num_taps : FIR length per polyphase arm - fixes XLA buffer allocation
21
+ # stride : decimation factor (== sps, typically 2 for T/2-spaced input)
22
+ # num_ch : MIMO butterfly width C - fixes matrix shapes in XLA IR
23
+ #
24
+ # Kernel-specific static variable:
25
+ # const_size : constellation size M for LMS/RLS - fixes the slicer argmin
26
+ # shape at trace time so XLA can compile the min-distance
27
+ # search without dynamic dispatch. CMA is blind (no slicer),
28
+ # so this variable is omitted from its key.
29
+ #
30
+ # All kernels share the same output convention:
31
+ # y_hat : (N_sym, C) complex64 - equalized symbols
32
+ # (transposed by _unpack_result_jax to (C, N_sym) before return)
33
+ # errors : (N_sym, C) complex64 - complex errors d - y
34
+ # w_hist : (N_sym, C, C, num_taps) complex64 - weight snapshots per symbol
35
+ #
36
+ # Performance notes:
37
+ # - jax.jit + XLA ahead-of-time compilation eliminates Python overhead for
38
+ # the scan loop; each symbol step is a single XLA op dispatch.
39
+ # - jax.vmap(get_win) vectorises window extraction across C channels.
40
+ # - jnp.einsum('ijt,jt->i') compiles to an optimised GEMV via OpenBLAS
41
+ # on CPU or cuBLAS on GPU.
42
+ # - lax.dynamic_slice emits a single gather op with runtime offset and
43
+ # compile-time slice size - no Python-level indexing overhead.
44
+ # - For GPU: lax.scan serialises the symbol loop on a single CUDA stream,
45
+ # limiting parallelism to the per-step GEMV. GPU is only beneficial
46
+ # for large MIMO widths (C >> 4) that saturate the cuBLAS kernel.
47
+ # Use backend='numba' for CPU-optimal throughput on typical SISO/2x2.
48
+
49
+ _JITTED_EQ: dict[tuple[Any, ...], Any] = {}
50
+
51
+
52
+ def _get_jax_lms(
53
+ num_taps, stride, const_size, num_ch, sq_side=0, sq_lev_min=0.0, sq_d_grid=1.0
54
+ ):
55
+ """JIT-compile and cache the sample-by-sample LMS butterfly scan.
56
+
57
+ Static closure variables (baked into the compiled kernel; a new cache
58
+ entry is created - not a retrace - when any of these change):
59
+
60
+ num_taps : FIR filter length per polyphase arm.
61
+ stride : decimation factor (== sps, typically 2 for T/2-spaced input).
62
+ const_size : constellation size M - fixes the slicer ``argmin`` shape at
63
+ trace time so XLA can compile it without dynamic dispatch.
64
+ num_ch : MIMO butterfly width C (number of input/output channels).
65
+ sq_side : int - 0 means O(M) slicer; >0 enables O(1) square-QAM slicer.
66
+ sq_lev_min, sq_d_grid : float - constellation level grid parameters.
67
+
68
+ Returns
69
+ -------
70
+ lms_scan : JIT-compiled callable
71
+ See the inner function for the call signature.
72
+ """
73
+ key = (
74
+ "lms",
75
+ num_taps,
76
+ stride,
77
+ const_size,
78
+ num_ch,
79
+ sq_side,
80
+ float(sq_lev_min),
81
+ float(sq_d_grid),
82
+ )
83
+ if key not in _JITTED_EQ:
84
+ jax, jnp, _ = _get_jax()
85
+
86
+ @jax.jit
87
+ def lms_scan(
88
+ x_input, training_padded, constellation, w_init, step_size, n_train
89
+ ):
90
+ # x_input : (C, N_pad) complex64
91
+ # training_padded : (C, N_sym) complex64
92
+ # constellation : (M,) complex64 - slicer lookup table
93
+ # w_init : (C, C, num_taps) complex64
94
+ # step_size : scalar float32
95
+ # n_train : scalar int32
96
+
97
+ def step(W, idx):
98
+ sample_idx = idx * stride
99
+
100
+ X_wins = jax.lax.dynamic_slice(
101
+ x_input, (0, sample_idx), (num_ch, num_taps)
102
+ )
103
+
104
+ _P = jax.lax.Precision.HIGHEST
105
+ y = jnp.einsum("ijt,jt->i", jnp.conj(W), X_wins, precision=_P) # (C,)
106
+
107
+ def slicer(ch_y):
108
+ if sq_side > 0: # static branch at trace time
109
+ ir = jnp.clip(
110
+ jnp.round((ch_y.real - sq_lev_min) / sq_d_grid).astype(
111
+ jnp.int32
112
+ ),
113
+ 0,
114
+ sq_side - 1,
115
+ )
116
+ ii = jnp.clip(
117
+ jnp.round((ch_y.imag - sq_lev_min) / sq_d_grid).astype(
118
+ jnp.int32
119
+ ),
120
+ 0,
121
+ sq_side - 1,
122
+ )
123
+ nr = sq_lev_min + ir.astype(jnp.float32) * jnp.float32(
124
+ sq_d_grid
125
+ )
126
+ ni = sq_lev_min + ii.astype(jnp.float32) * jnp.float32(
127
+ sq_d_grid
128
+ )
129
+ return jax.lax.complex(nr, ni)
130
+ else:
131
+ return constellation[
132
+ jnp.argmin(jnp.abs(ch_y - constellation) ** 2)
133
+ ]
134
+
135
+ dd = jax.vmap(slicer)(y) # (C,)
136
+ d = jnp.where(idx < n_train, training_padded[:, idx], dd)
137
+
138
+ e = d - y # (C,)
139
+
140
+ W_new = W + step_size * jnp.einsum("i,jt->ijt", jnp.conj(e), X_wins)
141
+ return W_new, (y, e, W_new)
142
+
143
+ n_sym = training_padded.shape[1]
144
+ W_final, (y_hat, errors, w_hist) = jax.lax.scan(
145
+ step, w_init, jnp.arange(n_sym)
146
+ )
147
+ return y_hat, errors, W_final, w_hist
148
+
149
+ _JITTED_EQ[key] = lms_scan
150
+ return _JITTED_EQ[key]
151
+
152
+
153
+ def _get_jax_rls(
154
+ num_taps, stride, const_size, num_ch, sq_side=0, sq_lev_min=0.0, sq_d_grid=1.0
155
+ ):
156
+ """JIT-compile and cache the sample-by-sample Leaky-RLS butterfly scan.
157
+
158
+ Static closure variables (same semantics as ``_get_jax_lms``):
159
+ num_taps, stride, const_size, num_ch.
160
+ sq_side, sq_lev_min, sq_d_grid: O(1) square-QAM slicer parameters.
161
+
162
+ Returns
163
+ -------
164
+ rls_scan : JIT-compiled callable
165
+ See the inner function for the call signature.
166
+ """
167
+ key = (
168
+ "rls",
169
+ num_taps,
170
+ stride,
171
+ const_size,
172
+ num_ch,
173
+ sq_side,
174
+ float(sq_lev_min),
175
+ float(sq_d_grid),
176
+ )
177
+ if key not in _JITTED_EQ:
178
+ jax, jnp, _ = _get_jax()
179
+
180
+ @jax.jit
181
+ def rls_scan(
182
+ x_input,
183
+ training_padded,
184
+ constellation,
185
+ w_init,
186
+ P_init,
187
+ lam,
188
+ n_train,
189
+ leakage,
190
+ n_update_halt,
191
+ ):
192
+ # Argument shapes and semantics
193
+ # ------------------------------
194
+ # x_input : (C, N_pad) complex64 - padded received samples
195
+ # training_padded : (C, N_sym) complex64 - reference symbols
196
+ # constellation : (M,) complex64 - slicer lookup table
197
+ # w_init : (C, C, num_taps) complex64 - initial butterfly weights
198
+ # P_init : (C*num_taps, C*num_taps) complex64 - initial inverse
199
+ # correlation matrix P = (1/delta) * I
200
+ # lam : scalar float32 - forgetting factor λ ∈ (0,1]
201
+ # n_train : scalar int32 - training/DD boundary (dynamic)
202
+ # leakage : scalar float32 - weight-decay coefficient γ ∈ [0,1):
203
+ # W <- (1-γ)W + k⊗ē each step; P update is unchanged.
204
+ # Decays null-subspace weights without inflating P eigenvalues.
205
+ # n_update_halt : scalar int32 - freeze W and P beyond this index;
206
+ # = n_sym - num_taps//2 to avoid zero-padding contamination
207
+ #
208
+ # lax.scan carry : (W, P)
209
+ # W : (C, C, num_taps) butterfly weight matrix
210
+ # P : (C*num_taps, C*num_taps) inverse input correlation matrix
211
+ # lax.scan xs : jnp.arange(N_sym)
212
+ # lax.scan output : y_hat (N_sym, C) equalized symbols
213
+ # errors (N_sym, C) complex errors d - y
214
+ # w_hist (N_sym, C, C, num_taps) weight snapshots
215
+ #
216
+ # Per-step equations (idx is the current symbol index):
217
+ # X_wins = [x_input[c, idx*stride : idx*stride+num_taps] for c] (C, T)
218
+ # x_bar = X_wins.flatten() (C*T,)
219
+ # y = einsum('ijt,jt->i', conj(W), X_wins) (C,)
220
+ # d = training[:, idx] if idx < n_train else slicer(y) (C,)
221
+ # e = d - y (C,)
222
+ # Px = P @ x_bar (C*T,)
223
+ # k = Px / (λ + x_bar^H @ Px) (C*T,) Kalman gain
224
+ # W <- (1-γ)W + k ⊗ conj(e) (if idx < n_update_halt)
225
+ # P = (P - outer(k, x_bar^H P)) / λ (if idx < n_update_halt)
226
+
227
+ _P = jax.lax.Precision.HIGHEST
228
+
229
+ def step(carry, idx):
230
+ W, P = carry
231
+ sample_idx = idx * stride
232
+
233
+ X_wins = jax.lax.dynamic_slice(
234
+ x_input, (0, sample_idx), (num_ch, num_taps)
235
+ )
236
+
237
+ y = jnp.einsum("ijt,jt->i", jnp.conj(W), X_wins, precision=_P)
238
+
239
+ def slicer(ch_y):
240
+ if sq_side > 0: # static branch at trace time
241
+ ir = jnp.clip(
242
+ jnp.round((ch_y.real - sq_lev_min) / sq_d_grid).astype(
243
+ jnp.int32
244
+ ),
245
+ 0,
246
+ sq_side - 1,
247
+ )
248
+ ii = jnp.clip(
249
+ jnp.round((ch_y.imag - sq_lev_min) / sq_d_grid).astype(
250
+ jnp.int32
251
+ ),
252
+ 0,
253
+ sq_side - 1,
254
+ )
255
+ nr = sq_lev_min + ir.astype(jnp.float32) * jnp.float32(
256
+ sq_d_grid
257
+ )
258
+ ni = sq_lev_min + ii.astype(jnp.float32) * jnp.float32(
259
+ sq_d_grid
260
+ )
261
+ return jax.lax.complex(nr, ni)
262
+ else:
263
+ return constellation[
264
+ jnp.argmin(jnp.abs(ch_y - constellation) ** 2)
265
+ ]
266
+
267
+ dd = jax.vmap(slicer)(y)
268
+ d = jnp.where(idx < n_train, training_padded[:, idx], dd)
269
+ e = d - y
270
+
271
+ x_bar = X_wins.flatten() # (C * num_taps,)
272
+
273
+ Px = jnp.matmul(P, x_bar, precision=_P)
274
+ denom = lam + jnp.real(jnp.dot(jnp.conj(x_bar), Px, precision=_P))
275
+ k = Px / denom
276
+
277
+ def w_update(w_row, err_val):
278
+ w_flat = w_row.flatten()
279
+ # Weight decay: suppress null-subspace taps exponentially.
280
+ # Adding γI directly to P would inflate P eigenvalues, making
281
+ # R_xx more singular and worsening AWGN amplification.
282
+ w_flat_new = (1.0 - leakage) * w_flat + k * jnp.conj(err_val)
283
+ return w_flat_new.reshape(num_ch, num_taps)
284
+
285
+ W_upd = jax.vmap(w_update)(W, e).astype(jnp.complex64)
286
+ # Riccati: exploit Hermitian symmetry P = P^H so (x^H P)[j] = conj((Px)[j]).
287
+ # outer(k, conj(Px)) == k ⊗ (x^H P); reuses Px to avoid a second O(N²) mat-vec.
288
+ P_upd = (P - jnp.outer(k, jnp.conj(Px))) / lam
289
+ # Hermitian re-symmetrization: P <- (P + Pᴴ)/2
290
+ # Prevents asymmetry drift from the 1/λ amplification.
291
+ P_upd = 0.5 * (P_upd + jnp.conj(P_upd).T)
292
+
293
+ # Early halt: freeze W and P once the sliding window begins
294
+ # overlapping the right zero-padding (last num_taps//2 symbols).
295
+ # The forward pass (y) continues so all output symbols are produced.
296
+ update_ok = idx < n_update_halt
297
+ W_new = jnp.where(update_ok, W_upd, W)
298
+ P_new = jnp.where(update_ok, P_upd, P)
299
+
300
+ return (W_new, P_new), (y, e, W_new)
301
+
302
+ n_sym = training_padded.shape[1]
303
+ (W_final, _), (y_hat, errors, w_hist) = jax.lax.scan(
304
+ step, (w_init, P_init), jnp.arange(n_sym)
305
+ )
306
+ return y_hat, errors, W_final, w_hist
307
+
308
+ _JITTED_EQ[key] = rls_scan
309
+ return _JITTED_EQ[key]
310
+
311
+
312
+ def _get_jax_lms_cpr(
313
+ num_taps,
314
+ stride,
315
+ const_size,
316
+ num_ch,
317
+ cpr_type,
318
+ bps_n,
319
+ bps_block_size,
320
+ bps_joint_channels,
321
+ cs_history_len,
322
+ symmetry=4,
323
+ sq_side=0,
324
+ sq_lev_min=0.0,
325
+ sq_d_grid=1.0,
326
+ ):
327
+ """JIT-compile and cache the LMS+CPR butterfly scan.
328
+
329
+ All CPR parameters are static closure variables (baked into the XLA graph
330
+ at trace time). A separate cache entry is created for each distinct
331
+ combination of these parameters.
332
+
333
+ cpr_type : "pll" or "bps"
334
+ bps_n : number of BPS test phases (ignored for cpr_type="pll")
335
+ cs_history_len : int - circular buffer depth for cycle-slip correction
336
+ sq_side, sq_lev_min, sq_d_grid : O(1) square-QAM slicer parameters
337
+ """
338
+ key = (
339
+ "lms_cpr",
340
+ num_taps,
341
+ stride,
342
+ const_size,
343
+ num_ch,
344
+ cpr_type,
345
+ bps_n,
346
+ bps_block_size,
347
+ bps_joint_channels,
348
+ cs_history_len,
349
+ int(symmetry),
350
+ sq_side,
351
+ float(sq_lev_min),
352
+ float(sq_d_grid),
353
+ )
354
+ if key not in _JITTED_EQ:
355
+ jax, jnp, _ = _get_jax()
356
+
357
+ H = cs_history_len
358
+ KB = bps_block_size # static closure: BPS window length
359
+
360
+ import math as _math
361
+
362
+ _quantum_static = jnp.float64(2.0 * _math.pi / symmetry)
363
+
364
+ _PREC = jax.lax.Precision.HIGHEST
365
+
366
+ @jax.jit
367
+ def lms_cpr_scan(
368
+ x_input,
369
+ training_padded,
370
+ constellation,
371
+ bps_phases_neg,
372
+ bps_angles,
373
+ w_init,
374
+ step_size,
375
+ n_train,
376
+ pll_mu,
377
+ pll_beta,
378
+ cs_threshold,
379
+ cs_enabled,
380
+ pll_phi_init,
381
+ pll_freq_init,
382
+ bps_buf_init,
383
+ bps_buf_ptr_init,
384
+ bps_prev4_init,
385
+ cs_buf_x_init,
386
+ cs_buf_y_init,
387
+ cs_buf_ptr_init,
388
+ ):
389
+ # x_input : (C, N_pad) complex64
390
+ # training_padded : (C, N_sym) complex64
391
+ # constellation : (M,) complex64
392
+ # bps_phases_neg : (B,) complex64 exp(-j*theta_k)
393
+ # bps_angles : (B,) float32 theta_k
394
+ # w_init : (C, C, T) complex64
395
+ # step_size : scalar float32
396
+ # n_train : scalar int32
397
+ # pll_mu, pll_beta: scalar float64
398
+ # cs_threshold : scalar float64
399
+ # cs_enabled : scalar bool
400
+ # pll_phi_init : (C,) float64 - warm-start PLL integrator (zeros -> cold)
401
+ # pll_freq_init : (C,) float64 - warm-start PLL frequency (zeros -> cold)
402
+ # bps_buf_init : (KB, C) complex64 - warm-start BPS buffer (zeros -> cold)
403
+ # bps_buf_ptr_init: scalar int32 - warm-start BPS buffer pointer
404
+ # bps_prev4_init : (C,) float64 - warm-start 4-fold unwrap state
405
+ # cs_buf_x_init : (C, H) float64 - warm-start cycle-slip symbol index
406
+ # cs_buf_y_init : (C, H) float64 - warm-start cycle-slip phase value
407
+ # cs_buf_ptr_init : (C,) int32 - warm-start cycle-slip write pointer
408
+ #
409
+ # lax.scan carry:
410
+ # W : (C, C, T) complex64
411
+ # pll_phi : (C,) float64
412
+ # pll_freq : (C,) float64
413
+ # bps_buf : (KB, C) complex64 - y_raw circular buffer
414
+ # bps_buf_ptr : scalar int32
415
+ # bps_prev4 : (C,) float64 - causal 4-fold unwrap state
416
+ # cs_buf_x : (C, H) float64 - symbol index
417
+ # cs_buf_y : (C, H) float64 - phase value
418
+ # cs_buf_ptr : (C,) int32 - write pointer
419
+ # bps_d2_slots : (B, KB, C) float64 - per-slot BPS sq-dist
420
+ # bps_metric : (B, C) float64 - running sum of slots
421
+
422
+ def _bps_d2(rotated):
423
+ # Min squared distance of `rotated` (..., complex) to the
424
+ # constellation, returned float64. Shared by the per-symbol
425
+ # new-slot update and the warm-start reconstruction so the
426
+ # distance formula has a single source of truth.
427
+ if sq_side > 0: # static branch at trace time
428
+ r_idx = jnp.clip(
429
+ jnp.round((rotated.real - sq_lev_min) / sq_d_grid).astype(
430
+ jnp.int32
431
+ ),
432
+ 0,
433
+ sq_side - 1,
434
+ )
435
+ i_idx = jnp.clip(
436
+ jnp.round((rotated.imag - sq_lev_min) / sq_d_grid).astype(
437
+ jnp.int32
438
+ ),
439
+ 0,
440
+ sq_side - 1,
441
+ )
442
+ r_near = sq_lev_min + r_idx.astype(jnp.float32) * jnp.float32(
443
+ sq_d_grid
444
+ )
445
+ i_near = sq_lev_min + i_idx.astype(jnp.float32) * jnp.float32(
446
+ sq_d_grid
447
+ )
448
+ d2 = (rotated.real - r_near) ** 2 + (rotated.imag - i_near) ** 2
449
+ return d2.astype(jnp.float64)
450
+ return jnp.min(
451
+ jnp.abs(rotated[..., None] - constellation) ** 2,
452
+ axis=-1,
453
+ ).astype(jnp.float64)
454
+
455
+ def step(carry, idx):
456
+ (
457
+ W,
458
+ pll_phi,
459
+ pll_freq,
460
+ bps_buf,
461
+ bps_buf_ptr,
462
+ bps_prev4,
463
+ bps_d2_slots,
464
+ bps_metric,
465
+ cs_buf_x,
466
+ cs_buf_y,
467
+ cs_buf_ptr,
468
+ ) = carry
469
+ sample_idx = idx * stride
470
+
471
+ X_wins = jax.lax.dynamic_slice(
472
+ x_input, (0, sample_idx), (num_ch, num_taps)
473
+ ) # (C, T)
474
+
475
+ y_raw = jnp.einsum(
476
+ "ijt,jt->i",
477
+ jnp.conj(W).astype(jnp.complex128),
478
+ X_wins.astype(jnp.complex128),
479
+ precision=_PREC,
480
+ ).astype(
481
+ jnp.complex64
482
+ ) # float64 accumulation -> complex64, matches Numba
483
+
484
+ # -- Phase estimation (static branch at trace time) --
485
+ if cpr_type == "pll":
486
+ phi_hat = pll_phi # (C,)
487
+ bps_buf_new = bps_buf
488
+ bps_buf_ptr_new = bps_buf_ptr
489
+ bps_prev4_new = bps_prev4
490
+ bps_d2_slots_new = bps_d2_slots
491
+ bps_metric_new = bps_metric
492
+ else:
493
+ # Fill BPS circular buffer with current y_raw
494
+ slot = jnp.int32(bps_buf_ptr % KB)
495
+ bps_buf_new = jax.lax.dynamic_update_slice(
496
+ bps_buf,
497
+ y_raw[None, :], # (1, C) - update one row
498
+ (slot, jnp.int32(0)),
499
+ ) # broadcast doesn't work for transpose, use (C, KB) layout instead
500
+ bps_buf_ptr_new = bps_buf_ptr + 1
501
+
502
+ # Incremental running sum (O(B*M)/symbol vs O(B*KB*M) for a
503
+ # full re-sum, matching the Numba bps_running_sum): score the
504
+ # NEW slot only, add it to the metric and subtract the evicted
505
+ # slot's stored distance. Not-yet-filled slots hold 0 (set at
506
+ # warm-start reconstruction), so the subtraction is exact and
507
+ # reproduces the old fill-mask semantics.
508
+ rotated_new = bps_phases_neg[:, None] * y_raw[None, :] # (B, C)
509
+ d2_new = _bps_d2(rotated_new) # (B, C) float64
510
+ old_d2 = jax.lax.dynamic_slice_in_dim(
511
+ bps_d2_slots, slot, 1, axis=1
512
+ )[:, 0, :] # (B, C)
513
+ metric = bps_metric + d2_new - old_d2 # (B, C) float64
514
+ bps_d2_slots_new = jax.lax.dynamic_update_slice_in_dim(
515
+ bps_d2_slots, d2_new[:, None, :], slot, axis=1
516
+ )
517
+ bps_metric_new = metric
518
+
519
+ if bps_joint_channels:
520
+ # Sum over channels too -> (B,); broadcast winner to all C
521
+ best_k = jnp.argmin(metric.sum(axis=-1)) # scalar
522
+ phi_raw = jnp.full(num_ch, bps_angles[best_k]) # (C,)
523
+ else:
524
+ best_k = jnp.argmin(metric, axis=0) # (C,)
525
+ phi_raw = bps_angles[best_k] # (C,)
526
+
527
+ # Causal 4-fold phase unwrap: equivalent to np.unwrap(phi*4)/4
528
+ raw4 = phi_raw.astype(jnp.float64) * jnp.float64(4.0)
529
+ diff4 = raw4 - bps_prev4
530
+ two_pi = jnp.float64(2.0 * 3.141592653589793)
531
+ diff4 = diff4 - jnp.round(diff4 / two_pi) * two_pi
532
+ bps_prev4_new = bps_prev4 + diff4
533
+ phi_hat = bps_prev4_new / jnp.float64(4.0) # (C,) unwrapped float64
534
+
535
+ # -- Cycle-slip correction ----------------------------
536
+ def correct_slip_ch(phi_h, buf_x_ch, buf_y_ch, ptr_ch):
537
+ # buf_x_ch is retained for call-site compat but unused here.
538
+ fill_cs = jnp.minimum(ptr_ch, H)
539
+ y_b = phi_h
540
+
541
+ mask = jnp.arange(H) < fill_cs
542
+ n_f = fill_cs.astype(jnp.float64)
543
+
544
+ # Sx and Sxx are exact closed-form constants - no cancellation.
545
+ Sx = n_f * (n_f - jnp.float64(1.0)) / jnp.float64(2.0)
546
+ Sxx = (
547
+ n_f
548
+ * (n_f - jnp.float64(1.0))
549
+ * (jnp.float64(2.0) * n_f - jnp.float64(1.0))
550
+ / jnp.float64(6.0)
551
+ )
552
+ denom = n_f * Sxx - Sx * Sx
553
+
554
+ # Sxy uses relative positions [0, fill_cs-1]. In the circular
555
+ # buffer the oldest entry is at slot ptr_ch%H and has position 0;
556
+ # slot j has relative position (j - ptr_ch%H + H) % H for a full
557
+ # window, or just j for a partial window (entries written 0..fill-1).
558
+ oldest_slot = ptr_ch % H
559
+ slots = jnp.arange(H)
560
+ rel_pos_full = (slots - oldest_slot + H) % H
561
+ rel_pos_partial = slots
562
+ rel_pos = jnp.where(
563
+ fill_cs >= H, rel_pos_full, rel_pos_partial
564
+ ).astype(jnp.float64)
565
+ Sy = jnp.where(mask, buf_y_ch, jnp.float64(0.0)).sum()
566
+ Sxy = jnp.where(mask, rel_pos * buf_y_ch, jnp.float64(0.0)).sum()
567
+
568
+ # Prediction target: one step past the newest entry (relative pos = fill_cs).
569
+ x_pred = n_f
570
+
571
+ safe_denom = jnp.where(
572
+ jnp.abs(denom) > 1e-20, denom, jnp.float64(1.0)
573
+ )
574
+ slope = jnp.where(
575
+ fill_cs >= 10,
576
+ (n_f * Sxy - Sx * Sy) / safe_denom,
577
+ jnp.float64(0.0),
578
+ )
579
+ intercept = jnp.where(
580
+ fill_cs >= 10,
581
+ (Sy - slope * Sx) / jnp.maximum(n_f, jnp.float64(1.0)),
582
+ Sy / jnp.maximum(n_f, jnp.float64(1.0)),
583
+ )
584
+ phi_exp_lin = slope * x_pred + intercept
585
+
586
+ last_pos = (ptr_ch - 1 + H) % H
587
+ phi_last = jax.lax.dynamic_index_in_dim(
588
+ buf_y_ch, last_pos, keepdims=False
589
+ )
590
+ phi_expected = jnp.where(fill_cs >= 10, phi_exp_lin, phi_last)
591
+ phi_expected = jnp.where(fill_cs == 0, y_b, phi_expected)
592
+
593
+ diff = y_b - phi_expected
594
+ k_slip = jnp.round(diff / _quantum_static)
595
+ should_correct = (
596
+ cs_enabled & (jnp.abs(diff) > cs_threshold) & (k_slip != 0)
597
+ )
598
+ phi_corr = jnp.where(
599
+ should_correct, y_b - k_slip * _quantum_static, y_b
600
+ )
601
+
602
+ write_pos = ptr_ch % H
603
+ buf_y_new = jax.lax.dynamic_update_slice(
604
+ buf_y_ch, phi_corr[None], [write_pos]
605
+ )
606
+ ptr_new = ptr_ch + 1
607
+ return phi_corr, buf_x_ch, buf_y_new, ptr_new
608
+
609
+ phi_corr, cs_buf_x_new, cs_buf_y_new, cs_buf_ptr_new = jax.vmap(
610
+ correct_slip_ch
611
+ )(phi_hat, cs_buf_x, cs_buf_y, cs_buf_ptr)
612
+
613
+ # Wrap to [-π, π] before casting to float32 for fast GPU exp
614
+ _two_pi_f64 = jnp.float64(2.0 * 3.141592653589793)
615
+ phi_wrapped = phi_corr - jnp.round(phi_corr / _two_pi_f64) * _two_pi_f64
616
+ _phi32 = phi_wrapped.astype(jnp.float32)
617
+ phasor = jnp.exp(_phi32 * jnp.array(-1j, dtype=jnp.complex64))
618
+ y_fin = y_raw * phasor # (C,)
619
+
620
+ def slicer(ch_y):
621
+ if sq_side > 0: # static branch at trace time
622
+ ir = jnp.clip(
623
+ jnp.round((ch_y.real - sq_lev_min) / sq_d_grid).astype(
624
+ jnp.int32
625
+ ),
626
+ 0,
627
+ sq_side - 1,
628
+ )
629
+ ii = jnp.clip(
630
+ jnp.round((ch_y.imag - sq_lev_min) / sq_d_grid).astype(
631
+ jnp.int32
632
+ ),
633
+ 0,
634
+ sq_side - 1,
635
+ )
636
+ nr = sq_lev_min + ir.astype(jnp.float32) * jnp.float32(
637
+ sq_d_grid
638
+ )
639
+ ni = sq_lev_min + ii.astype(jnp.float32) * jnp.float32(
640
+ sq_d_grid
641
+ )
642
+ return jax.lax.complex(nr, ni)
643
+ else:
644
+ return constellation[
645
+ jnp.argmin(jnp.abs(ch_y - constellation) ** 2)
646
+ ]
647
+
648
+ dd = jax.vmap(slicer)(y_fin)
649
+ d = jnp.where(idx < n_train, training_padded[:, idx], dd)
650
+ e_clean = d - y_fin # (C,)
651
+
652
+ if cpr_type == "pll":
653
+ e_ph = y_fin.imag * d.real - y_fin.real * d.imag # (C,)
654
+ if bps_joint_channels: # static branch - shared LO
655
+ e_ph = jnp.full(num_ch, e_ph.mean())
656
+ pll_phi_new = pll_phi + pll_mu * e_ph + pll_freq
657
+ pll_freq_new = pll_freq + pll_beta * e_ph
658
+ else:
659
+ pll_phi_new = pll_phi
660
+ pll_freq_new = pll_freq
661
+
662
+ phasor_inv = jnp.exp(_phi32 * jnp.array(1j, dtype=jnp.complex64))
663
+ e_eq = e_clean * phasor_inv # (C,)
664
+
665
+ W_new = W + step_size * jnp.einsum("i,jt->ijt", jnp.conj(e_eq), X_wins)
666
+
667
+ carry_new = (
668
+ W_new,
669
+ pll_phi_new,
670
+ pll_freq_new,
671
+ bps_buf_new,
672
+ bps_buf_ptr_new,
673
+ bps_prev4_new,
674
+ bps_d2_slots_new,
675
+ bps_metric_new,
676
+ cs_buf_x_new,
677
+ cs_buf_y_new,
678
+ cs_buf_ptr_new,
679
+ )
680
+ return carry_new, (y_fin, e_clean, W_new, phi_corr)
681
+
682
+ n_sym = training_padded.shape[1]
683
+ # Warm-start: reconstruct per-slot BPS distances + running metric from
684
+ # the buffer (single source of truth via _bps_d2). Cold start (zero
685
+ # buffer, ptr=0) yields an all-zero metric through the fill mask; the
686
+ # mask also stores 0 for unfilled slots so the in-loop eviction stays
687
+ # exact.
688
+ _fill0 = jnp.minimum(bps_buf_ptr_init, KB)
689
+ _rot0 = bps_phases_neg[:, None, None] * bps_buf_init[None, :, :] # (B,KB,C)
690
+ _mask0 = jnp.arange(KB)[None, :, None] < _fill0
691
+ bps_d2_slots_init = jnp.where(_mask0, _bps_d2(_rot0), jnp.float64(0.0))
692
+ bps_metric_init = bps_d2_slots_init.sum(axis=1) # (B, C) float64
693
+ init_carry = (
694
+ w_init,
695
+ pll_phi_init,
696
+ pll_freq_init,
697
+ bps_buf_init,
698
+ bps_buf_ptr_init,
699
+ bps_prev4_init,
700
+ bps_d2_slots_init,
701
+ bps_metric_init,
702
+ cs_buf_x_init,
703
+ cs_buf_y_init,
704
+ cs_buf_ptr_init,
705
+ )
706
+ (
707
+ (
708
+ W_final,
709
+ pll_phi_f,
710
+ pll_freq_f,
711
+ bps_buf_f,
712
+ bps_buf_ptr_f,
713
+ bps_prev4_f,
714
+ _bps_d2_slots_f,
715
+ _bps_metric_f,
716
+ cs_buf_x_f,
717
+ cs_buf_y_f,
718
+ cs_buf_ptr_f,
719
+ ),
720
+ (y_hat, errors, w_hist, phi_traj),
721
+ ) = jax.lax.scan(step, init_carry, jnp.arange(n_sym))
722
+ return (
723
+ y_hat,
724
+ errors,
725
+ W_final,
726
+ w_hist,
727
+ phi_traj,
728
+ pll_phi_f,
729
+ pll_freq_f,
730
+ bps_buf_f,
731
+ bps_buf_ptr_f,
732
+ bps_prev4_f,
733
+ cs_buf_x_f,
734
+ cs_buf_y_f,
735
+ cs_buf_ptr_f,
736
+ )
737
+
738
+ _JITTED_EQ[key] = lms_cpr_scan
739
+ return _JITTED_EQ[key]
740
+
741
+
742
+ def _get_jax_rls_cpr(
743
+ num_taps,
744
+ stride,
745
+ const_size,
746
+ num_ch,
747
+ cpr_type,
748
+ bps_n,
749
+ bps_block_size,
750
+ bps_joint_channels,
751
+ cs_history_len,
752
+ symmetry=4,
753
+ sq_side=0,
754
+ sq_lev_min=0.0,
755
+ sq_d_grid=1.0,
756
+ ):
757
+ """JIT-compile and cache the RLS+CPR butterfly scan.
758
+
759
+ Combines the Leaky-RLS Riccati update with an inline CPR tracker.
760
+ Static parameters are identical to ``_get_jax_lms_cpr``.
761
+ sq_side, sq_lev_min, sq_d_grid : O(1) square-QAM slicer parameters.
762
+ """
763
+ key = (
764
+ "rls_cpr",
765
+ num_taps,
766
+ stride,
767
+ const_size,
768
+ num_ch,
769
+ cpr_type,
770
+ bps_n,
771
+ bps_block_size,
772
+ bps_joint_channels,
773
+ cs_history_len,
774
+ int(symmetry),
775
+ sq_side,
776
+ float(sq_lev_min),
777
+ float(sq_d_grid),
778
+ )
779
+ if key not in _JITTED_EQ:
780
+ jax, jnp, _ = _get_jax()
781
+
782
+ H = cs_history_len
783
+ KB = bps_block_size
784
+ import math as _math
785
+
786
+ _quantum_static = jnp.float64(2.0 * _math.pi / symmetry)
787
+
788
+ _PREC = jax.lax.Precision.HIGHEST
789
+
790
+ @jax.jit
791
+ def rls_cpr_scan(
792
+ x_input,
793
+ training_padded,
794
+ constellation,
795
+ bps_phases_neg,
796
+ bps_angles,
797
+ w_init,
798
+ P_init,
799
+ lam,
800
+ n_train,
801
+ leakage,
802
+ n_update_halt,
803
+ pll_mu,
804
+ pll_beta,
805
+ cs_threshold,
806
+ cs_enabled,
807
+ pll_phi_init,
808
+ pll_freq_init,
809
+ bps_buf_init,
810
+ bps_buf_ptr_init,
811
+ bps_prev4_init,
812
+ cs_buf_x_init,
813
+ cs_buf_y_init,
814
+ cs_buf_ptr_init,
815
+ ):
816
+ def _bps_d2(rotated):
817
+ # Min squared distance of `rotated` (..., complex) to the
818
+ # constellation, returned float64. Shared by the per-symbol
819
+ # new-slot update and the warm-start reconstruction so the
820
+ # distance formula has a single source of truth.
821
+ if sq_side > 0: # static branch at trace time
822
+ r_idx = jnp.clip(
823
+ jnp.round((rotated.real - sq_lev_min) / sq_d_grid).astype(
824
+ jnp.int32
825
+ ),
826
+ 0,
827
+ sq_side - 1,
828
+ )
829
+ i_idx = jnp.clip(
830
+ jnp.round((rotated.imag - sq_lev_min) / sq_d_grid).astype(
831
+ jnp.int32
832
+ ),
833
+ 0,
834
+ sq_side - 1,
835
+ )
836
+ r_near = sq_lev_min + r_idx.astype(jnp.float32) * jnp.float32(
837
+ sq_d_grid
838
+ )
839
+ i_near = sq_lev_min + i_idx.astype(jnp.float32) * jnp.float32(
840
+ sq_d_grid
841
+ )
842
+ d2 = (rotated.real - r_near) ** 2 + (rotated.imag - i_near) ** 2
843
+ return d2.astype(jnp.float64)
844
+ return jnp.min(
845
+ jnp.abs(rotated[..., None] - constellation) ** 2,
846
+ axis=-1,
847
+ ).astype(jnp.float64)
848
+
849
+ def step(carry, idx):
850
+ (
851
+ W,
852
+ P,
853
+ pll_phi,
854
+ pll_freq,
855
+ bps_buf,
856
+ bps_buf_ptr,
857
+ bps_prev4,
858
+ bps_d2_slots,
859
+ bps_metric,
860
+ cs_buf_x,
861
+ cs_buf_y,
862
+ cs_buf_ptr,
863
+ ) = carry
864
+ sample_idx = idx * stride
865
+
866
+ X_wins = jax.lax.dynamic_slice(
867
+ x_input, (0, sample_idx), (num_ch, num_taps)
868
+ )
869
+ y_raw = jnp.einsum(
870
+ "ijt,jt->i",
871
+ jnp.conj(W).astype(jnp.complex128),
872
+ X_wins.astype(jnp.complex128),
873
+ precision=_PREC,
874
+ ).astype(
875
+ jnp.complex64
876
+ ) # float64 accumulation -> complex64, matches Numba
877
+
878
+ if cpr_type == "pll":
879
+ phi_hat = pll_phi
880
+ bps_buf_new = bps_buf
881
+ bps_buf_ptr_new = bps_buf_ptr
882
+ bps_prev4_new = bps_prev4
883
+ bps_d2_slots_new = bps_d2_slots
884
+ bps_metric_new = bps_metric
885
+ else:
886
+ slot = jnp.int32(bps_buf_ptr % KB)
887
+ bps_buf_new = jax.lax.dynamic_update_slice(
888
+ bps_buf,
889
+ y_raw[None, :],
890
+ (slot, jnp.int32(0)),
891
+ )
892
+ bps_buf_ptr_new = bps_buf_ptr + 1
893
+
894
+ # Incremental running sum (O(B*M)/symbol vs O(B*KB*M) for a
895
+ # full re-sum, matching the Numba bps_running_sum): score the
896
+ # NEW slot only, add it to the metric and subtract the evicted
897
+ # slot's stored distance. Not-yet-filled slots hold 0 (set at
898
+ # warm-start reconstruction), so the subtraction is exact and
899
+ # reproduces the old fill-mask semantics.
900
+ rotated_new = bps_phases_neg[:, None] * y_raw[None, :] # (B, C)
901
+ d2_new = _bps_d2(rotated_new) # (B, C) float64
902
+ old_d2 = jax.lax.dynamic_slice_in_dim(
903
+ bps_d2_slots, slot, 1, axis=1
904
+ )[:, 0, :] # (B, C)
905
+ metric = bps_metric + d2_new - old_d2 # (B, C) float64
906
+ bps_d2_slots_new = jax.lax.dynamic_update_slice_in_dim(
907
+ bps_d2_slots, d2_new[:, None, :], slot, axis=1
908
+ )
909
+ bps_metric_new = metric
910
+
911
+ if bps_joint_channels:
912
+ best_k = jnp.argmin(metric.sum(axis=-1))
913
+ phi_raw = jnp.full(num_ch, bps_angles[best_k])
914
+ else:
915
+ best_k = jnp.argmin(metric, axis=0)
916
+ phi_raw = bps_angles[best_k]
917
+
918
+ # Causal 4-fold phase unwrap
919
+ raw4 = phi_raw.astype(jnp.float64) * jnp.float64(4.0)
920
+ diff4 = raw4 - bps_prev4
921
+ two_pi = jnp.float64(2.0 * 3.141592653589793)
922
+ diff4 = diff4 - jnp.round(diff4 / two_pi) * two_pi
923
+ bps_prev4_new = bps_prev4 + diff4
924
+ phi_hat = bps_prev4_new / jnp.float64(4.0)
925
+
926
+ def correct_slip_ch(phi_h, buf_x_ch, buf_y_ch, ptr_ch):
927
+ # buf_x_ch is retained for call-site compat but unused here.
928
+ fill_cs = jnp.minimum(ptr_ch, H)
929
+ y_b = phi_h
930
+ mask = jnp.arange(H) < fill_cs
931
+ n_f = fill_cs.astype(jnp.float64)
932
+
933
+ # Sx and Sxx are exact closed-form constants - no cancellation.
934
+ Sx = n_f * (n_f - jnp.float64(1.0)) / jnp.float64(2.0)
935
+ Sxx = (
936
+ n_f
937
+ * (n_f - jnp.float64(1.0))
938
+ * (jnp.float64(2.0) * n_f - jnp.float64(1.0))
939
+ / jnp.float64(6.0)
940
+ )
941
+ denom = n_f * Sxx - Sx * Sx
942
+
943
+ # Sxy uses relative positions [0, fill_cs-1] derived from the
944
+ # circular buffer layout (oldest slot = ptr_ch % H -> position 0).
945
+ oldest_slot = ptr_ch % H
946
+ slots = jnp.arange(H)
947
+ rel_pos_full = (slots - oldest_slot + H) % H
948
+ rel_pos = jnp.where(fill_cs >= H, rel_pos_full, slots).astype(
949
+ jnp.float64
950
+ )
951
+ Sy = jnp.where(mask, buf_y_ch, jnp.float64(0.0)).sum()
952
+ Sxy = jnp.where(mask, rel_pos * buf_y_ch, jnp.float64(0.0)).sum()
953
+
954
+ # Prediction target: one step past newest (relative pos = fill_cs).
955
+ x_pred = n_f
956
+
957
+ safe_denom = jnp.where(
958
+ jnp.abs(denom) > 1e-20, denom, jnp.float64(1.0)
959
+ )
960
+ slope = jnp.where(
961
+ fill_cs >= 10,
962
+ (n_f * Sxy - Sx * Sy) / safe_denom,
963
+ jnp.float64(0.0),
964
+ )
965
+ intercept = jnp.where(
966
+ fill_cs >= 10,
967
+ (Sy - slope * Sx) / jnp.maximum(n_f, jnp.float64(1.0)),
968
+ Sy / jnp.maximum(n_f, jnp.float64(1.0)),
969
+ )
970
+ phi_exp_lin = slope * x_pred + intercept
971
+ last_pos = (ptr_ch - 1 + H) % H
972
+ phi_last = jax.lax.dynamic_index_in_dim(
973
+ buf_y_ch, last_pos, keepdims=False
974
+ )
975
+ phi_expected = jnp.where(fill_cs >= 10, phi_exp_lin, phi_last)
976
+ phi_expected = jnp.where(fill_cs == 0, y_b, phi_expected)
977
+ diff = y_b - phi_expected
978
+ k_slip = jnp.round(diff / _quantum_static)
979
+ should_correct = (
980
+ cs_enabled & (jnp.abs(diff) > cs_threshold) & (k_slip != 0)
981
+ )
982
+ phi_corr = jnp.where(
983
+ should_correct, y_b - k_slip * _quantum_static, y_b
984
+ )
985
+ write_pos = ptr_ch % H
986
+ buf_y_new = jax.lax.dynamic_update_slice(
987
+ buf_y_ch, phi_corr[None], [write_pos]
988
+ )
989
+ ptr_new = ptr_ch + 1
990
+ return phi_corr, buf_x_ch, buf_y_new, ptr_new
991
+
992
+ phi_corr, cs_buf_x_new, cs_buf_y_new, cs_buf_ptr_new = jax.vmap(
993
+ correct_slip_ch
994
+ )(phi_hat, cs_buf_x, cs_buf_y, cs_buf_ptr)
995
+
996
+ # Wrap to [-π, π] before casting to float32 for fast GPU exp
997
+ _two_pi_f64 = jnp.float64(2.0 * 3.141592653589793)
998
+ phi_wrapped = phi_corr - jnp.round(phi_corr / _two_pi_f64) * _two_pi_f64
999
+ _phi32 = phi_wrapped.astype(jnp.float32)
1000
+ phasor = jnp.exp(_phi32 * jnp.array(-1j, dtype=jnp.complex64))
1001
+ y_fin = y_raw * phasor
1002
+
1003
+ def slicer(ch_y):
1004
+ if sq_side > 0: # static branch at trace time
1005
+ ir = jnp.clip(
1006
+ jnp.round((ch_y.real - sq_lev_min) / sq_d_grid).astype(
1007
+ jnp.int32
1008
+ ),
1009
+ 0,
1010
+ sq_side - 1,
1011
+ )
1012
+ ii = jnp.clip(
1013
+ jnp.round((ch_y.imag - sq_lev_min) / sq_d_grid).astype(
1014
+ jnp.int32
1015
+ ),
1016
+ 0,
1017
+ sq_side - 1,
1018
+ )
1019
+ nr = sq_lev_min + ir.astype(jnp.float32) * jnp.float32(
1020
+ sq_d_grid
1021
+ )
1022
+ ni = sq_lev_min + ii.astype(jnp.float32) * jnp.float32(
1023
+ sq_d_grid
1024
+ )
1025
+ return jax.lax.complex(nr, ni)
1026
+ else:
1027
+ return constellation[
1028
+ jnp.argmin(jnp.abs(ch_y - constellation) ** 2)
1029
+ ]
1030
+
1031
+ dd = jax.vmap(slicer)(y_fin)
1032
+ d = jnp.where(idx < n_train, training_padded[:, idx], dd)
1033
+ e_clean = d - y_fin
1034
+
1035
+ if cpr_type == "pll":
1036
+ e_ph = y_fin.imag * d.real - y_fin.real * d.imag
1037
+ if bps_joint_channels: # static branch - shared LO
1038
+ e_ph = jnp.full(num_ch, e_ph.mean())
1039
+ pll_phi_new = pll_phi + pll_mu * e_ph + pll_freq
1040
+ pll_freq_new = pll_freq + pll_beta * e_ph
1041
+ else:
1042
+ pll_phi_new = pll_phi
1043
+ pll_freq_new = pll_freq
1044
+
1045
+ phasor_inv = jnp.exp(_phi32 * jnp.array(1j, dtype=jnp.complex64))
1046
+ e_eq = e_clean * phasor_inv
1047
+
1048
+ x_bar = X_wins.flatten()
1049
+ Px = jnp.matmul(P, x_bar, precision=_PREC)
1050
+ denom_k = lam + jnp.real(jnp.dot(jnp.conj(x_bar), Px, precision=_PREC))
1051
+ k_gain = Px / denom_k
1052
+
1053
+ def w_update(w_row, err_val):
1054
+ w_flat = w_row.flatten()
1055
+ w_flat_new = (1.0 - leakage) * w_flat + k_gain * jnp.conj(err_val)
1056
+ return w_flat_new.reshape(num_ch, num_taps)
1057
+
1058
+ W_upd = jax.vmap(w_update)(W, e_eq).astype(jnp.complex64)
1059
+ P_upd = (P - jnp.outer(k_gain, jnp.conj(Px))) / lam
1060
+ P_upd = 0.5 * (P_upd + jnp.conj(P_upd).T)
1061
+
1062
+ update_ok = idx < n_update_halt
1063
+ W_new = jnp.where(update_ok, W_upd, W)
1064
+ P_new = jnp.where(update_ok, P_upd, P)
1065
+
1066
+ carry_new = (
1067
+ W_new,
1068
+ P_new,
1069
+ pll_phi_new,
1070
+ pll_freq_new,
1071
+ bps_buf_new,
1072
+ bps_buf_ptr_new,
1073
+ bps_prev4_new,
1074
+ bps_d2_slots_new,
1075
+ bps_metric_new,
1076
+ cs_buf_x_new,
1077
+ cs_buf_y_new,
1078
+ cs_buf_ptr_new,
1079
+ )
1080
+ return carry_new, (y_fin, e_clean, W_new, phi_corr)
1081
+
1082
+ n_sym = training_padded.shape[1]
1083
+ # Warm-start: reconstruct per-slot BPS distances + running metric from
1084
+ # the buffer (single source of truth via _bps_d2). Cold start (zero
1085
+ # buffer, ptr=0) yields an all-zero metric through the fill mask; the
1086
+ # mask also stores 0 for unfilled slots so the in-loop eviction stays
1087
+ # exact.
1088
+ _fill0 = jnp.minimum(bps_buf_ptr_init, KB)
1089
+ _rot0 = bps_phases_neg[:, None, None] * bps_buf_init[None, :, :] # (B,KB,C)
1090
+ _mask0 = jnp.arange(KB)[None, :, None] < _fill0
1091
+ bps_d2_slots_init = jnp.where(_mask0, _bps_d2(_rot0), jnp.float64(0.0))
1092
+ bps_metric_init = bps_d2_slots_init.sum(axis=1) # (B, C) float64
1093
+ init_carry = (
1094
+ w_init,
1095
+ P_init,
1096
+ pll_phi_init,
1097
+ pll_freq_init,
1098
+ bps_buf_init,
1099
+ bps_buf_ptr_init,
1100
+ bps_prev4_init,
1101
+ bps_d2_slots_init,
1102
+ bps_metric_init,
1103
+ cs_buf_x_init,
1104
+ cs_buf_y_init,
1105
+ cs_buf_ptr_init,
1106
+ )
1107
+ (
1108
+ (
1109
+ W_final,
1110
+ _,
1111
+ pll_phi_f,
1112
+ pll_freq_f,
1113
+ bps_buf_f,
1114
+ bps_buf_ptr_f,
1115
+ bps_prev4_f,
1116
+ _bps_d2_slots_f,
1117
+ _bps_metric_f,
1118
+ cs_buf_x_f,
1119
+ cs_buf_y_f,
1120
+ cs_buf_ptr_f,
1121
+ ),
1122
+ (y_hat, errors, w_hist, phi_traj),
1123
+ ) = jax.lax.scan(step, init_carry, jnp.arange(n_sym))
1124
+ return (
1125
+ y_hat,
1126
+ errors,
1127
+ W_final,
1128
+ w_hist,
1129
+ phi_traj,
1130
+ pll_phi_f,
1131
+ pll_freq_f,
1132
+ bps_buf_f,
1133
+ bps_buf_ptr_f,
1134
+ bps_prev4_f,
1135
+ cs_buf_x_f,
1136
+ cs_buf_y_f,
1137
+ cs_buf_ptr_f,
1138
+ )
1139
+
1140
+ _JITTED_EQ[key] = rls_cpr_scan
1141
+ return _JITTED_EQ[key]
1142
+
1143
+
1144
+ def _get_jax_cma(num_taps, stride, num_ch):
1145
+ """JIT-compile and cache the sample-by-sample CMA butterfly scan.
1146
+
1147
+ Static closure variables: num_taps, stride, num_ch (same as LMS/RLS).
1148
+ No constellation required - CMA is a blind algorithm.
1149
+
1150
+ Returns
1151
+ -------
1152
+ cma_scan : JIT-compiled callable
1153
+ See the inner function for the call signature.
1154
+ """
1155
+ key = ("cma", num_taps, stride, num_ch)
1156
+ if key not in _JITTED_EQ:
1157
+ jax, jnp, _ = _get_jax()
1158
+
1159
+ # n_sym (arg 4) is static: lax.scan requires a compile-time iteration
1160
+ # count. JAX's bounded LRU cache retraces when n_sym changes, keeping
1161
+ # _JITTED_EQ from growing without bound.
1162
+ @functools.partial(jax.jit, static_argnums=(4,))
1163
+ def cma_scan(x_input, w_init, step_size, r2, n_sym):
1164
+ # Argument shapes and semantics
1165
+ # ------------------------------
1166
+ # x_input : (C, N_pad) complex64 - padded received samples
1167
+ # w_init : (C, C, num_taps) complex64 - initial butterfly weights
1168
+ # step_size : scalar float32 - fixed gradient step μ (no NLMS; non-convex surface)
1169
+ # r2 : scalar float32 - Godard dispersion radius R² = E[|s|⁴] / E[|s|²]
1170
+ # n_sym : int (static) - total symbol count; fixes scan iteration count
1171
+ #
1172
+ # lax.scan carry : W (C, C, num_taps)
1173
+ # lax.scan xs : jnp.arange(n_sym)
1174
+ # lax.scan output : y_hat (n_sym, C)
1175
+ # errors (n_sym, C) Godard errors y*(|y|²-R²)
1176
+ # w_hist (n_sym, C, C, num_taps)
1177
+ #
1178
+ # Per-step gradient descent (Godard criterion):
1179
+ # y = einsum('ijt,jt->i', conj(W), X_wins) (C,)
1180
+ # e = y * (real(y * conj(y)) - R²) (C,) CMA error
1181
+ # W -= μ * einsum('i,jt->ijt', conj(e), X_wins) gradient step
1182
+ # Note: real() is required to prevent imaginary leakage from
1183
+ # floating-point noise in |y|² from causing parasitic phase rotation.
1184
+ _P = jax.lax.Precision.HIGHEST
1185
+
1186
+ def step(W, idx):
1187
+ sample_idx = idx * stride
1188
+
1189
+ X_wins = jax.lax.dynamic_slice(
1190
+ x_input, (0, sample_idx), (num_ch, num_taps)
1191
+ )
1192
+ y = jnp.einsum("ijt,jt->i", jnp.conj(W), X_wins, precision=_P)
1193
+
1194
+ # CMA error: e_i = y_i * (|y_i|^2 - R2)
1195
+ # jnp.real enforces strict real-valued modulus: floating-point noise
1196
+ # in y*conj(y) would otherwise inject imaginary components, causing
1197
+ # a parasitic phase rotation into the gradient via multiplication by y.
1198
+ e = y * (jnp.real(y * jnp.conj(y)) - r2)
1199
+
1200
+ W_new = W - step_size * jnp.einsum("i,jt->ijt", jnp.conj(e), X_wins)
1201
+ return W_new, (y, e, W_new)
1202
+
1203
+ W_final, (y_hat, errors, w_hist) = jax.lax.scan(
1204
+ step, w_init, jnp.arange(n_sym)
1205
+ )
1206
+ return y_hat, errors, W_final, w_hist
1207
+
1208
+ _JITTED_EQ[key] = cma_scan
1209
+ return _JITTED_EQ[key]
1210
+
1211
+
1212
+ def _get_jax_rde(num_taps, stride, num_radii, num_ch):
1213
+ """JIT-compile and cache the sample-by-sample RDE butterfly scan.
1214
+
1215
+ RDE (Radius Directed Equalizer) extends CMA by replacing the single
1216
+ Godard radius with per-symbol radius selection from a precomputed set
1217
+ of unique constellation magnitudes. This provides correct blind
1218
+ convergence on multi-ring constellations such as 16-QAM and 64-QAM.
1219
+
1220
+ Static closure variables (baked into the compiled kernel):
1221
+
1222
+ num_taps : FIR filter length per polyphase arm.
1223
+ stride : decimation factor (sps, typically 2 for T/2-spaced input).
1224
+ num_radii : number of unique radii K - fixes the argmin shape at trace
1225
+ time so XLA can compile without dynamic dispatch.
1226
+ num_ch : MIMO butterfly width C.
1227
+
1228
+ Returns
1229
+ -------
1230
+ rde_scan : JIT-compiled callable
1231
+ See the inner function for the call signature.
1232
+ """
1233
+ key = ("rde", num_taps, stride, num_radii, num_ch)
1234
+ if key not in _JITTED_EQ:
1235
+ jax, jnp, _ = _get_jax()
1236
+
1237
+ @functools.partial(jax.jit, static_argnums=(4,))
1238
+ def rde_scan(x_input, w_init, step_size, radii, n_sym):
1239
+ # Argument shapes and semantics
1240
+ # ------------------------------
1241
+ # x_input : (C, N_pad) complex64 - padded received samples
1242
+ # w_init : (C, C, num_taps) complex64 - initial butterfly weights
1243
+ # step_size : scalar float32 - fixed gradient step μ
1244
+ # radii : (K,) float32 - unique constellation radii, sorted
1245
+ # n_sym : int (static) - total symbol count; fixes scan iteration count
1246
+ #
1247
+ # lax.scan carry : W (C, C, num_taps)
1248
+ # lax.scan xs : jnp.arange(n_sym)
1249
+ # lax.scan output : y_hat (n_sym, C)
1250
+ # errors (n_sym, C) RDE errors y*(|y|²-R_d²)
1251
+ # w_hist (n_sym, C, C, num_taps)
1252
+ #
1253
+ # Per-step RDE gradient:
1254
+ # y = einsum('ijt,jt->i', conj(W), X_wins) (C,)
1255
+ # abs_y = sqrt(real(y*conj(y))) (C,)
1256
+ # R_d = radii[argmin(|radii-abs_y|)] (C,) nearest ring
1257
+ # e = y * (real(y*conj(y)) - R_d²) (C,)
1258
+ # W -= μ * einsum('i,jt->ijt', conj(e), X_wins)
1259
+ _P = jax.lax.Precision.HIGHEST
1260
+
1261
+ def step(W, idx):
1262
+ sample_idx = idx * stride
1263
+
1264
+ X_wins = jax.lax.dynamic_slice(
1265
+ x_input, (0, sample_idx), (num_ch, num_taps)
1266
+ )
1267
+ y = jnp.einsum("ijt,jt->i", jnp.conj(W), X_wins, precision=_P)
1268
+
1269
+ abs_y2 = jnp.real(y * jnp.conj(y)) # (C,) strict real |y|²
1270
+ abs_y = jnp.sqrt(abs_y2) # (C,) |y|
1271
+
1272
+ # (C, K) distance table; argmin over K gives nearest radius index
1273
+ dist = jnp.abs(abs_y[:, None] - radii[None, :]) # (C, K)
1274
+ rd = radii[jnp.argmin(dist, axis=1)] # (C,)
1275
+
1276
+ e = y * (abs_y2 - rd**2)
1277
+
1278
+ W_new = W - step_size * jnp.einsum("i,jt->ijt", jnp.conj(e), X_wins)
1279
+ return W_new, (y, e, W_new)
1280
+
1281
+ W_final, (y_hat, errors, w_hist) = jax.lax.scan(
1282
+ step, w_init, jnp.arange(n_sym)
1283
+ )
1284
+ return y_hat, errors, W_final, w_hist
1285
+
1286
+ _JITTED_EQ[key] = rde_scan
1287
+ return _JITTED_EQ[key]
1288
+
1289
+
1290
+ # -----------------------------------------------------------------------------
1291
+ # BLOCK-UPDATE EQUALIZERS (update_mode='block' - time-domain)
1292
+ # -----------------------------------------------------------------------------
1293
+ #
1294
+ # Block-update LMS/CMA/RDE freeze the butterfly weights over a chunk of ``D``
1295
+ # symbols, accumulate one aggregated gradient, and apply a single weight update
1296
+ # per chunk. This replaces the per-symbol weight dependency with one matrix
1297
+ # product per chunk, which XLA (JAX) and CuPy execute on the wide units - the
1298
+ # per-symbol ``lax.scan`` is launch-overhead-bound and cannot occupy a GPU.
1299
+ #
1300
+ # All variants share the **unified subtractive update**
1301
+ # W -= mu * sum_d conj(E[d]) ⊗ X[d]
1302
+ # (matching the pilot-aided kernels): LMS uses ``E = y - d`` (algebraically
1303
+ # identical to the additive ``W += mu·conj(d-y)·X``); CMA/RDE use the Godard /
1304
+ # ring error; pilot positions invert to ``E = y - pilot_ref``. Per CLAUDE.md
1305
+ # the two einsums run at ``Precision.HIGHEST`` to force true FP32 (no TF32).
1306
+
1307
+
1308
+ def _jax_block_core(
1309
+ jax,
1310
+ jnp,
1311
+ x_input,
1312
+ W_init,
1313
+ mu,
1314
+ error_fn,
1315
+ aux_xs,
1316
+ *,
1317
+ n_sym,
1318
+ D,
1319
+ stride,
1320
+ num_taps,
1321
+ num_ch,
1322
+ ):
1323
+ """Run a chunked block-update butterfly scan under an active jit trace.
1324
+
1325
+ Builds the per-symbol strided regressor windows, reshapes into
1326
+ ``n_chunks`` chunks of ``D`` symbols (zero-padding the final partial
1327
+ chunk and masking its gradient contribution), and scans over chunks with
1328
+ frozen weights. ``error_fn(Y_chunk, aux_chunk) -> E_chunk`` (both
1329
+ ``(D, num_ch)``) supplies the variant-specific error; the unified
1330
+ subtractive update is applied once per chunk.
1331
+
1332
+ Returns ``(y_hat (n_sym, C), errors (n_sym, C), W_final (C, C, T))``.
1333
+ """
1334
+ _P = jax.lax.Precision.HIGHEST
1335
+ n_chunks = (n_sym + D - 1) // D
1336
+ n_pad = n_chunks * D
1337
+
1338
+ # Per-symbol strided window gather -> (n_pad, C, num_taps). Padded symbols
1339
+ # (>= n_sym) index past the real signal; clamp to stay in-bounds - their
1340
+ # rows are masked out of the gradient and trimmed from the outputs.
1341
+ sym_ids = jnp.arange(n_pad)
1342
+ tap_ids = jnp.arange(num_taps)
1343
+ idx = sym_ids[:, None] * stride + tap_ids[None, :] # (n_pad, num_taps)
1344
+ idx = jnp.clip(idx, 0, x_input.shape[1] - 1)
1345
+ X_all = jnp.transpose(x_input[:, idx], (1, 0, 2)) # (n_pad, C, T)
1346
+ X_chunks = X_all.reshape(n_chunks, D, num_ch, num_taps)
1347
+
1348
+ valid = (sym_ids < n_sym).reshape(n_chunks, D) # (n_chunks, D)
1349
+
1350
+ def body(W, xs):
1351
+ X_chunk, valid_chunk, aux = xs
1352
+ Y = jnp.einsum("ijt,djt->di", jnp.conj(W), X_chunk, precision=_P) # (D,C)
1353
+ E = error_fn(Y, aux) # (D, C)
1354
+ E_masked = jnp.where(valid_chunk[:, None], E, jnp.zeros_like(E))
1355
+ grad = jnp.einsum("di,djt->ijt", jnp.conj(E_masked), X_chunk, precision=_P)
1356
+ # Accumulate at HIGHEST precision but keep complex64 weight storage
1357
+ # (CLAUDE.md): an x64-enabled session can otherwise promote the error to
1358
+ # complex128 and break the scan carry dtype.
1359
+ W_new = (W - mu * grad).astype(W.dtype)
1360
+ return W_new, (Y, E)
1361
+
1362
+ # unroll a few chunks per scan step to amortise XLA loop-control overhead -
1363
+ # the per-chunk matmuls are small, so loop dispatch is a real cost at the
1364
+ # default block_len.
1365
+ W_final, (Y_chunks, E_chunks) = jax.lax.scan(
1366
+ body, W_init, (X_chunks, valid, aux_xs), unroll=4
1367
+ )
1368
+ y_hat = Y_chunks.reshape(n_pad, num_ch)[:n_sym]
1369
+ errors = E_chunks.reshape(n_pad, num_ch)[:n_sym]
1370
+ return y_hat, errors, W_final
1371
+
1372
+
1373
+ def _jax_block_pilot_aux(
1374
+ jax, jnp, n_sym, D, num_ch, n_chunks, has_pilots, pref, pmask, blind_fn
1375
+ ):
1376
+ """Build the per-chunk pilot aux and wrap a blind error with a pilot override.
1377
+
1378
+ ``blind_fn(Y) -> E_blind`` (both ``(D, C)``). When ``has_pilots`` is True,
1379
+ masked positions invert to the LMS residual ``Y - pref`` (subtractive
1380
+ update, matching the per-symbol pilot-aided kernels); otherwise the aux is
1381
+ a dummy per-chunk index and the blind error is used everywhere.
1382
+
1383
+ Returns ``(aux_xs, error_fn)`` ready for ``_jax_block_core``.
1384
+ """
1385
+ n_pad = n_chunks * D
1386
+ if has_pilots:
1387
+ pref_T = jnp.transpose(pref) # (n_sym, C)
1388
+ pref_pad = jnp.zeros((n_pad, num_ch), pref_T.dtype).at[:n_sym].set(pref_T)
1389
+ pmask_pad = (
1390
+ jnp.zeros((n_pad,), jnp.bool_).at[:n_sym].set(pmask.astype(jnp.bool_))
1391
+ )
1392
+ aux_xs = (
1393
+ pref_pad.reshape(n_chunks, D, num_ch),
1394
+ pmask_pad.reshape(n_chunks, D),
1395
+ )
1396
+
1397
+ def error_fn(Y, aux):
1398
+ pref_c, pmask_c = aux
1399
+ return jnp.where(pmask_c[:, None], Y - pref_c, blind_fn(Y))
1400
+
1401
+ else:
1402
+ aux_xs = jnp.arange(n_chunks)
1403
+
1404
+ def error_fn(Y, aux):
1405
+ return blind_fn(Y)
1406
+
1407
+ return aux_xs, error_fn
1408
+
1409
+
1410
+ def _get_jax_lms_block(
1411
+ num_taps,
1412
+ stride,
1413
+ const_size,
1414
+ num_ch,
1415
+ n_sym,
1416
+ D,
1417
+ sq_side=0,
1418
+ sq_lev_min=0.0,
1419
+ sq_d_grid=1.0,
1420
+ ):
1421
+ """JIT-compile and cache the block-update LMS butterfly scan.
1422
+
1423
+ Static closure variables mirror ``_get_jax_lms`` plus ``n_sym`` and the
1424
+ update block length ``D`` (both fix the chunk count / scan length at trace
1425
+ time). Training/DD switch is applied elementwise within each chunk.
1426
+ """
1427
+ key = (
1428
+ "lms_block",
1429
+ num_taps,
1430
+ stride,
1431
+ const_size,
1432
+ num_ch,
1433
+ n_sym,
1434
+ D,
1435
+ sq_side,
1436
+ float(sq_lev_min),
1437
+ float(sq_d_grid),
1438
+ )
1439
+ if key not in _JITTED_EQ:
1440
+ jax, jnp, _ = _get_jax()
1441
+ n_chunks = (n_sym + D - 1) // D
1442
+ n_pad = n_chunks * D
1443
+
1444
+ def _slice_block(Y, constellation):
1445
+ if sq_side > 0: # static branch - O(1) square-QAM slicer
1446
+ ir = jnp.clip(
1447
+ jnp.round((Y.real - sq_lev_min) / sq_d_grid).astype(jnp.int32),
1448
+ 0,
1449
+ sq_side - 1,
1450
+ )
1451
+ ii = jnp.clip(
1452
+ jnp.round((Y.imag - sq_lev_min) / sq_d_grid).astype(jnp.int32),
1453
+ 0,
1454
+ sq_side - 1,
1455
+ )
1456
+ nr = sq_lev_min + ir.astype(jnp.float32) * jnp.float32(sq_d_grid)
1457
+ ni = sq_lev_min + ii.astype(jnp.float32) * jnp.float32(sq_d_grid)
1458
+ return jax.lax.complex(nr, ni)
1459
+ d2 = jnp.abs(Y[..., None] - constellation) ** 2 # (D, C, M)
1460
+ return constellation[jnp.argmin(d2, axis=-1)]
1461
+
1462
+ @jax.jit
1463
+ def run(x_input, training_padded, constellation, W_init, mu, n_train):
1464
+ # training_padded: (C, n_sym) -> per-symbol (n_pad, C) chunks
1465
+ train_T = jnp.transpose(training_padded) # (n_sym, C)
1466
+ train_pad = jnp.zeros((n_pad, num_ch), train_T.dtype)
1467
+ train_pad = train_pad.at[:n_sym].set(train_T)
1468
+ train_chunks = train_pad.reshape(n_chunks, D, num_ch)
1469
+ gidx_chunks = jnp.arange(n_pad).reshape(n_chunks, D)
1470
+
1471
+ def error_fn(Y, aux):
1472
+ train_chunk, gidx = aux # (D, C), (D,)
1473
+ dd = _slice_block(Y, constellation)
1474
+ d = jnp.where((gidx < n_train)[:, None], train_chunk, dd)
1475
+ return Y - d
1476
+
1477
+ y_hat, errors, W_final = _jax_block_core(
1478
+ jax,
1479
+ jnp,
1480
+ x_input,
1481
+ W_init,
1482
+ mu,
1483
+ error_fn,
1484
+ (train_chunks, gidx_chunks),
1485
+ n_sym=n_sym,
1486
+ D=D,
1487
+ stride=stride,
1488
+ num_taps=num_taps,
1489
+ num_ch=num_ch,
1490
+ )
1491
+ return y_hat, errors, W_final, jnp.zeros((1,), jnp.complex64)
1492
+
1493
+ _JITTED_EQ[key] = run
1494
+ return _JITTED_EQ[key]
1495
+
1496
+
1497
+ def _get_jax_cma_block(num_taps, stride, num_ch, n_sym, D, has_pilots=False):
1498
+ """JIT-compile and cache the block-update CMA butterfly scan.
1499
+
1500
+ Blind Godard error per chunk; when ``has_pilots`` the error inverts to the
1501
+ LMS residual ``y - pilot_ref`` at masked positions (subtractive update,
1502
+ matching ``_get_numba_pa_cma``).
1503
+ """
1504
+ key = ("cma_block", num_taps, stride, num_ch, n_sym, D, has_pilots)
1505
+ if key not in _JITTED_EQ:
1506
+ jax, jnp, _ = _get_jax()
1507
+ n_chunks = (n_sym + D - 1) // D
1508
+
1509
+ @jax.jit
1510
+ def run(x_input, W_init, mu, r2, pref, pmask):
1511
+ aux_xs, error_fn = _jax_block_pilot_aux(
1512
+ jax,
1513
+ jnp,
1514
+ n_sym,
1515
+ D,
1516
+ num_ch,
1517
+ n_chunks,
1518
+ has_pilots,
1519
+ pref,
1520
+ pmask,
1521
+ lambda Y: Y * (jnp.real(Y * jnp.conj(Y)) - r2),
1522
+ )
1523
+ y_hat, errors, W_final = _jax_block_core(
1524
+ jax,
1525
+ jnp,
1526
+ x_input,
1527
+ W_init,
1528
+ mu,
1529
+ error_fn,
1530
+ aux_xs,
1531
+ n_sym=n_sym,
1532
+ D=D,
1533
+ stride=stride,
1534
+ num_taps=num_taps,
1535
+ num_ch=num_ch,
1536
+ )
1537
+ return y_hat, errors, W_final, jnp.zeros((1,), jnp.complex64)
1538
+
1539
+ _JITTED_EQ[key] = run
1540
+ return _JITTED_EQ[key]
1541
+
1542
+
1543
+ def _get_jax_rde_block(num_taps, stride, num_radii, num_ch, n_sym, D, has_pilots=False):
1544
+ """JIT-compile and cache the block-update RDE butterfly scan.
1545
+
1546
+ Per-symbol nearest-ring radial error per chunk; pilot positions invert to
1547
+ the LMS residual as in CMA.
1548
+ """
1549
+ key = ("rde_block", num_taps, stride, num_radii, num_ch, n_sym, D, has_pilots)
1550
+ if key not in _JITTED_EQ:
1551
+ jax, jnp, _ = _get_jax()
1552
+ n_chunks = (n_sym + D - 1) // D
1553
+
1554
+ @jax.jit
1555
+ def run(x_input, W_init, mu, radii, pref, pmask):
1556
+ def blind_fn(Y):
1557
+ abs_y2 = jnp.real(Y * jnp.conj(Y)) # (D, C)
1558
+ abs_y = jnp.sqrt(abs_y2)
1559
+ dist = jnp.abs(abs_y[..., None] - radii[None, None, :]) # (D, C, K)
1560
+ rd = radii[jnp.argmin(dist, axis=-1)] # (D, C)
1561
+ return Y * (abs_y2 - rd**2)
1562
+
1563
+ aux_xs, error_fn = _jax_block_pilot_aux(
1564
+ jax, jnp, n_sym, D, num_ch, n_chunks, has_pilots, pref, pmask, blind_fn
1565
+ )
1566
+ y_hat, errors, W_final = _jax_block_core(
1567
+ jax,
1568
+ jnp,
1569
+ x_input,
1570
+ W_init,
1571
+ mu,
1572
+ error_fn,
1573
+ aux_xs,
1574
+ n_sym=n_sym,
1575
+ D=D,
1576
+ stride=stride,
1577
+ num_taps=num_taps,
1578
+ num_ch=num_ch,
1579
+ )
1580
+ return y_hat, errors, W_final, jnp.zeros((1,), jnp.complex64)
1581
+
1582
+ _JITTED_EQ[key] = run
1583
+ return _JITTED_EQ[key]
1584
+
1585
+
1586
+ def _get_jax_pa_cma(num_taps: int, stride: int, num_ch: int):
1587
+ """JIT-compile and cache the JAX pilot-aided CMA butterfly scan.
1588
+
1589
+ Hybrid CMA scan using ``jax.lax.scan``. At pilot positions
1590
+ (``pilot_mask[t] == True``) the error is the standard LMS residual
1591
+ ``pilot_ref[t] - y``; at data positions the Godard CMA error
1592
+ ``y * (|y|² - R²)`` is used. The switch is XLA-branchless via
1593
+ ``jnp.where``, keeping the scan body shape-static for efficient
1594
+ compilation.
1595
+
1596
+ Parameters
1597
+ ----------
1598
+ num_taps : int
1599
+ stride : int - samples per symbol (sps)
1600
+ num_ch : int - number of MIMO channels C
1601
+
1602
+ Returns
1603
+ -------
1604
+ pa_cma_scan : jax.jit-compiled callable
1605
+ ``(x_input, w_init, step_size, r2, pilot_ref, pilot_mask, n_sym)
1606
+ -> (y_all, e_all, W_final, w_hist)``
1607
+ where ``pilot_ref`` has shape ``(n_sym, C)`` and ``pilot_mask``
1608
+ has shape ``(n_sym,)`` bool.
1609
+ """
1610
+ key = ("pa_cma", num_taps, stride, num_ch)
1611
+ if key not in _JITTED_EQ:
1612
+ jax, jnp, _ = _get_jax()
1613
+ assert jax is not None
1614
+ assert jnp is not None
1615
+
1616
+ @functools.partial(jax.jit, static_argnums=(6,))
1617
+ def pa_cma_scan(
1618
+ x_input, # (C, N_pad) complex64 - padded received samples
1619
+ w_init, # (C, C, T) complex64 - initial butterfly weights
1620
+ step_size, # () float32
1621
+ r2, # () float32 - Godard R² = E[|s|⁴]/E[|s|²]
1622
+ pilot_ref, # (n_sym, C) complex64 - known symbols; 0 at data
1623
+ pilot_mask, # (n_sym,) bool - True at pilot/preamble positions
1624
+ n_sym, # int (static) - total symbol count
1625
+ ):
1626
+ _P = jax.lax.Precision.HIGHEST
1627
+
1628
+ def step(W, xs_t):
1629
+ idx, p_ref, p_mask = xs_t # (): int, (C,): cplx, (): bool
1630
+ X_wins = jax.lax.dynamic_slice(
1631
+ x_input, (0, idx * stride), (num_ch, num_taps)
1632
+ ) # (C, T)
1633
+ y = jnp.einsum("ijt,jt->i", jnp.conj(W), X_wins, precision=_P) # (C,)
1634
+
1635
+ abs_y2 = jnp.real(y * jnp.conj(y)) # strict real |y|²
1636
+ e_blind = y * (abs_y2 - r2) # Godard CMA
1637
+ e_da = (
1638
+ y - p_ref
1639
+ ) # pilot LMS (inverted to match blind subtractive update)
1640
+ e = jnp.where(p_mask, e_da, e_blind) # (C,) branchless
1641
+
1642
+ W_new = W - step_size * jnp.einsum("i,jt->ijt", jnp.conj(e), X_wins)
1643
+ return W_new, (y, e, W_new)
1644
+
1645
+ xs = (jnp.arange(n_sym), pilot_ref, pilot_mask)
1646
+ W_final, (y_all, e_all, wh_all) = jax.lax.scan(step, w_init, xs)
1647
+ return y_all, e_all, W_final, wh_all
1648
+
1649
+ _JITTED_EQ[key] = pa_cma_scan # type: ignore[index]
1650
+ return _JITTED_EQ[key] # type: ignore[index]
1651
+
1652
+
1653
+ def _get_jax_pa_rde(num_taps: int, stride: int, num_radii: int, num_ch: int):
1654
+ """JIT-compile and cache the JAX pilot-aided RDE butterfly scan.
1655
+
1656
+ Hybrid RDE scan using ``jax.lax.scan``. At pilot positions the error
1657
+ is the LMS residual ``pilot_ref[t] - y``; at data positions the
1658
+ ring-directed RDE error ``y * (|y|² - R_d²)`` is used, where ``R_d``
1659
+ is the nearest constellation ring radius. The switch is branchless.
1660
+
1661
+ Parameters
1662
+ ----------
1663
+ num_taps : int
1664
+ stride : int - samples per symbol (sps)
1665
+ num_radii : int - number of unique constellation ring radii K
1666
+ num_ch : int - number of MIMO channels C
1667
+
1668
+ Returns
1669
+ -------
1670
+ pa_rde_scan : jax.jit-compiled callable
1671
+ ``(x_input, w_init, step_size, radii, pilot_ref, pilot_mask, n_sym)
1672
+ -> (y_all, e_all, W_final, w_hist)``
1673
+ where ``pilot_ref`` has shape ``(n_sym, C)`` and ``pilot_mask``
1674
+ has shape ``(n_sym,)`` bool.
1675
+ """
1676
+ key = ("pa_rde", num_taps, stride, num_radii, num_ch)
1677
+ if key not in _JITTED_EQ:
1678
+ jax, jnp, _ = _get_jax()
1679
+ assert jax is not None
1680
+ assert jnp is not None
1681
+
1682
+ @functools.partial(jax.jit, static_argnums=(6,))
1683
+ def pa_rde_scan(
1684
+ x_input, # (C, N_pad) complex64 - padded received samples
1685
+ w_init, # (C, C, T) complex64 - initial butterfly weights
1686
+ step_size, # () float32
1687
+ radii, # (K,) float32 - unique |c| constellation radii, sorted
1688
+ pilot_ref, # (n_sym, C) complex64 - known symbols; 0 at data
1689
+ pilot_mask, # (n_sym,) bool - True at pilot/preamble positions
1690
+ n_sym, # int (static) - total symbol count
1691
+ ):
1692
+ _P = jax.lax.Precision.HIGHEST
1693
+
1694
+ def step(W, xs_t):
1695
+ idx, p_ref, p_mask = xs_t # (): int, (C,): cplx, (): bool
1696
+ X_wins = jax.lax.dynamic_slice(
1697
+ x_input, (0, idx * stride), (num_ch, num_taps)
1698
+ ) # (C, T)
1699
+ y = jnp.einsum("ijt,jt->i", jnp.conj(W), X_wins, precision=_P) # (C,)
1700
+
1701
+ abs_y2 = jnp.real(y * jnp.conj(y)) # strict real |y|²
1702
+ abs_y = jnp.sqrt(abs_y2) # (C,)
1703
+ dist = jnp.abs(abs_y[:, None] - radii[None, :]) # (C, K) broadcast
1704
+ rd = radii[jnp.argmin(dist, axis=1)] # (C,) nearest radius
1705
+
1706
+ e_blind = y * (abs_y2 - rd**2) # RDE ring-directed
1707
+ e_da = (
1708
+ y - p_ref
1709
+ ) # pilot LMS (inverted to match blind subtractive update)
1710
+ e = jnp.where(p_mask, e_da, e_blind) # (C,) branchless
1711
+
1712
+ W_new = W - step_size * jnp.einsum("i,jt->ijt", jnp.conj(e), X_wins)
1713
+ return W_new, (y, e, W_new)
1714
+
1715
+ xs = (jnp.arange(n_sym), pilot_ref, pilot_mask)
1716
+ W_final, (y_all, e_all, wh_all) = jax.lax.scan(step, w_init, xs)
1717
+ return y_all, e_all, W_final, wh_all
1718
+
1719
+ _JITTED_EQ[key] = pa_rde_scan # type: ignore[index]
1720
+ return _JITTED_EQ[key] # type: ignore[index]