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.
- commkit/__init__.py +74 -0
- commkit/_cuda/__init__.py +321 -0
- commkit/_cuda/compiler.py +88 -0
- commkit/_cuda/src/bps_min_d2.cu +104 -0
- commkit/_cuda/src/cs_block.cu +119 -0
- commkit/_cuda/src/selftest.cu +14 -0
- commkit/analysis/__init__.py +55 -0
- commkit/analysis/_common.py +236 -0
- commkit/analysis/allan.py +108 -0
- commkit/analysis/drift.py +213 -0
- commkit/analysis/interferometry.py +887 -0
- commkit/analysis/linewidth.py +480 -0
- commkit/analysis/trajectory.py +91 -0
- commkit/backend.py +507 -0
- commkit/coding/__init__.py +23 -0
- commkit/coding/base.py +17 -0
- commkit/coding/bch.py +6 -0
- commkit/coding/convolutional.py +7 -0
- commkit/coding/crc.py +7 -0
- commkit/coding/galois.py +8 -0
- commkit/coding/hamming.py +6 -0
- commkit/coding/interleaving.py +7 -0
- commkit/coding/ldpc.py +8 -0
- commkit/coding/polar.py +8 -0
- commkit/coding/ratematch.py +6 -0
- commkit/coding/reed_solomon.py +6 -0
- commkit/coding/turbo.py +8 -0
- commkit/core/__init__.py +32 -0
- commkit/core/frame.py +992 -0
- commkit/core/generation.py +581 -0
- commkit/core/signal.py +725 -0
- commkit/equalization/__init__.py +49 -0
- commkit/equalization/_block.py +1855 -0
- commkit/equalization/_common.py +606 -0
- commkit/equalization/_kernels_jax.py +1720 -0
- commkit/equalization/_kernels_numba.py +1704 -0
- commkit/equalization/blind.py +223 -0
- commkit/equalization/linear.py +365 -0
- commkit/equalization/polarization.py +790 -0
- commkit/equalization/result.py +191 -0
- commkit/equalization/sequential.py +2805 -0
- commkit/filtering.py +1120 -0
- commkit/frequency.py +1191 -0
- commkit/helpers.py +489 -0
- commkit/impairments/__init__.py +43 -0
- commkit/impairments/channel/__init__.py +20 -0
- commkit/impairments/channel/linear.py +310 -0
- commkit/impairments/channel/nonlinear.py +11 -0
- commkit/impairments/frontend.py +229 -0
- commkit/impairments/noise.py +105 -0
- commkit/impairments/source.py +219 -0
- commkit/io.py +308 -0
- commkit/logger.py +103 -0
- commkit/mapping/__init__.py +46 -0
- commkit/mapping/bits.py +240 -0
- commkit/mapping/constellation.py +153 -0
- commkit/mapping/gray.py +429 -0
- commkit/mapping/llr.py +253 -0
- commkit/mapping/shaping.py +218 -0
- commkit/metrics.py +949 -0
- commkit/multirate.py +476 -0
- commkit/plotting/__init__.py +78 -0
- commkit/plotting/analysis.py +627 -0
- commkit/plotting/constellation.py +483 -0
- commkit/plotting/equalizer.py +390 -0
- commkit/plotting/eye.py +388 -0
- commkit/plotting/spectral.py +575 -0
- commkit/plotting/sync.py +953 -0
- commkit/plotting/theme.py +203 -0
- commkit/plotting/waveform.py +200 -0
- commkit/py.typed +0 -0
- commkit/recovery/__init__.py +51 -0
- commkit/recovery/bps.py +337 -0
- commkit/recovery/corrections.py +751 -0
- commkit/recovery/pilots.py +803 -0
- commkit/recovery/pll.py +482 -0
- commkit/recovery/tikhonov.py +424 -0
- commkit/recovery/viterbi_viterbi.py +227 -0
- commkit/spectral.py +560 -0
- commkit/timing.py +841 -0
- commkit-1.0.0.dist-info/METADATA +145 -0
- commkit-1.0.0.dist-info/RECORD +84 -0
- commkit-1.0.0.dist-info/WHEEL +4 -0
- 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]
|