rf-compute 0.1.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.
rf_compute/__init__.py ADDED
@@ -0,0 +1,38 @@
1
+ """rf-compute: wave-domain computation kernel — package init."""
2
+
3
+ from .rf_compute import (
4
+ Operator, AirCompOperator, LatticeAirCompOperator, FadingAirCompOperator,
5
+ OTAAggregationOperator, ConvolutionOperator, InversionOperator, ReservoirOperator,
6
+ WaveComputeKernel,
7
+ boxcar, differencer, matched, hilbert,
8
+ )
9
+ from . import lattice
10
+ from . import coefficients
11
+ from . import ota_fl
12
+ from .lattice import (
13
+ mod_lattice, encode, decode, channel, run_trial, monte_carlo,
14
+ fading_trial, fading_scoreline,
15
+ )
16
+ from .coefficients import (
17
+ mmse_alpha, computation_rate, norm_bound, select_coefficients,
18
+ fading_gains, lll_reduce,
19
+ )
20
+ from .ota_fl import (
21
+ make_federated_data, gradient_spread, ota_aggregate, aggregation_quality,
22
+ )
23
+
24
+ __version__ = "0.1.0"
25
+ __all__ = [
26
+ 'Operator', 'AirCompOperator', 'LatticeAirCompOperator',
27
+ 'FadingAirCompOperator', 'OTAAggregationOperator',
28
+ 'ConvolutionOperator', 'InversionOperator', 'ReservoirOperator',
29
+ 'WaveComputeKernel',
30
+ 'boxcar', 'differencer', 'matched', 'hilbert',
31
+ 'lattice', 'coefficients', 'ota_fl',
32
+ 'mod_lattice', 'encode', 'decode', 'channel', 'run_trial', 'monte_carlo',
33
+ 'fading_trial', 'fading_scoreline',
34
+ 'mmse_alpha', 'computation_rate', 'norm_bound', 'select_coefficients',
35
+ 'fading_gains', 'lll_reduce',
36
+ 'make_federated_data', 'gradient_spread', 'ota_aggregate', 'aggregation_quality',
37
+ '__version__',
38
+ ]
@@ -0,0 +1,328 @@
1
+ """
2
+ Coefficient selection for compute-and-forward — the fading-channel tier
3
+ ========================================================================
4
+
5
+ Tier 1.5 (rf_compute/lattice.py) SETS the channel gains to integers
6
+ (h_i = a_i), which makes the compute-and-forward equation exact — but it
7
+ hides the problem the coefficient vector exists to solve: on a real channel
8
+ the gains are FADING (real-valued, not integers), and the receiver must
9
+ CHOOSE an integer vector a whose combination sum(a_i * w_i) mod L it wants
10
+ to decode.
11
+
12
+ This module implements the selection machinery:
13
+
14
+ computation_rate(h, a, snr_db) — the Nazer/Gastpar rate for one a
15
+ mmse_alpha(h, a, snr_db) — the optimal receiver scaling alpha
16
+ select_coefficients(h, snr_db, ..) — exhaustive (norm bound) / LLL / rounded
17
+ fading_gains(num_nodes, rng) — a real fading realization
18
+ lll_reduce(basis) — Lenstra-Lenstra-Lovasz reduction
19
+
20
+ The theory in one paragraph
21
+ ---------------------------
22
+ A real channel y = sum_i h_i x_i + z. The receiver scales by alpha and
23
+ replays the shared dithers weighted at the effective gains alpha*h_i:
24
+
25
+ y' = alpha*y + sum_i alpha*h_i*d_i
26
+ = sum_i alpha*h_i*v_i + L*(integers) + alpha*z
27
+
28
+ Mod-L reduction kills the lattice components; what remains is the real
29
+ number sum(alpha*h_i*w_i) + alpha*z, which the decoder rounds to the
30
+ nearest integer — the combination sum(a_i*w_i) mod L, PROVIDED alpha*h_i
31
+ is close to the integers a_i. The mismatch is the self-noise:
32
+
33
+ Z_eff(alpha, a) = alpha^2 sigma^2 + P * sum_i (alpha*h_i - a_i)^2
34
+
35
+ with P the per-node transmit power (L^2/12, the uniform-cell convention)
36
+ and sigma^2 = P / SNR. The computation rate is
37
+
38
+ R(alpha, a) = 1/2 log2^+ ( P / Z_eff )
39
+
40
+ maximized over alpha and a. Two facts make this tractable:
41
+
42
+ 1. The optimal scaling is the MMSE choice (closed form):
43
+ alpha* = SNR * (h^T a) / (1 + SNR * ||h||^2)
44
+ 2. Substituting alpha* gives the shortest-lattice-vector (SLV) problem
45
+ minimize D(a) = a^T (I + SNR * h h^T)^{-1} a over integer a != 0
46
+
47
+ Sanity check that pins the normalization: N = 1, h = 1, a = 1 gives
48
+ D = 1/(1+SNR) and R = 1/2 log2(1 + SNR) — the AWGN capacity of a real
49
+ channel, exactly. (At high SNR the integer alignment alpha*h ~ a can only
50
+ be perfect when a is parallel to h — the residual sum_i (alpha*h_i - a_i)^2
51
+ is the self-noise floor. This is why the real-channel toy never reaches the
52
+ cooperative bound for N >= 2 unless h is integer-aligned, and why the
53
+ literature works over COMPLEX channels (two real dimensions; Gaussian-
54
+ integer lattices align far better) — Liu & Ling 2016.)
55
+
56
+ Search methods (the literature's lineage, in order)
57
+ ---------------------------------------------------
58
+ - exhaustive within the norm bound ||a|| <= sqrt(1 + SNR ||h||^2):
59
+ Nazer & Gastpar 2011 (the theorem's own bound; exact, exponential in N)
60
+ - LLL lattice reduction: approximate the SLV instance in polynomial time.
61
+ The lineage: Sahraei & Gastpar 2014 (exact polynomial algorithm),
62
+ Liu & Ling 2016 (complex channels, efficient integer search)
63
+ - rounded: a = round(alpha*h) — the naive nearest-integer heuristic,
64
+ kept as the baseline (it is what you get if you never optimize)
65
+
66
+ No hardware, no secrets — pure NumPy.
67
+ """
68
+
69
+ from __future__ import annotations
70
+
71
+ import itertools
72
+
73
+ import numpy as np
74
+
75
+ __all__ = [
76
+ 'mmse_alpha', 'computation_rate', 'norm_bound', 'select_coefficients',
77
+ 'fading_gains', 'lll_reduce',
78
+ ]
79
+
80
+
81
+ # ─────────────────────────────────────────────────────────────────────────────
82
+ # The rate and the scaling
83
+ # ─────────────────────────────────────────────────────────────────────────────
84
+
85
+ def mmse_alpha(h, a, snr_db):
86
+ """
87
+ The MMSE receiver scaling: alpha* = SNR * (h^T a) / (1 + SNR ||h||^2).
88
+
89
+ Minimizes the effective decoder noise Z_eff(alpha) = alpha^2 sigma^2 +
90
+ P * sum_i (alpha*h_i - a_i)^2 (Nazer & Gastpar 2011). Returns 0.0 when
91
+ h^T a <= 0 (no positive scaling helps — the rate is 0).
92
+ """
93
+ h = np.asarray(h, dtype=np.float64)
94
+ a = np.asarray(a, dtype=np.float64)
95
+ snr = 10 ** (snr_db / 10)
96
+ num = snr * float(h @ a)
97
+ if num <= 0:
98
+ return 0.0
99
+ return num / (1.0 + snr * float(h @ h))
100
+
101
+
102
+ def computation_rate(h, a, snr_db, alpha=None):
103
+ """
104
+ The compute-and-forward computation rate for coefficient vector a:
105
+ R = 1/2 log2^+ ( P / Z_eff ), Z_eff = alpha^2 sigma^2 + P ||alpha h - a||^2
106
+ reported in bits per real channel use, with alpha = MMSE unless given.
107
+
108
+ Normalization: per-node transmit power P, sigma^2 = P / SNR — the
109
+ per-user power convention of the AirComp literature (Huang & Burr 2017).
110
+ N = 1, h = 1, a = 1 returns exactly 1/2 log2(1 + SNR).
111
+ """
112
+ h = np.asarray(h, dtype=np.float64)
113
+ a = np.asarray(a, dtype=np.float64)
114
+ snr = 10 ** (snr_db / 10)
115
+ if alpha is None:
116
+ alpha = mmse_alpha(h, a, snr_db)
117
+ if alpha <= 0:
118
+ return 0.0
119
+ # Z_eff / P = alpha^2 / SNR + ||alpha h - a||^2
120
+ d = alpha ** 2 / snr + float(np.sum((alpha * h - a) ** 2))
121
+ if d <= 0:
122
+ return 0.0
123
+ return 0.5 * max(0.0, float(np.log2(1.0 / d)))
124
+
125
+
126
+ def norm_bound(h, snr_db):
127
+ """
128
+ The Nazer/Gastpar search bound: any rate-maximizing a satisfies
129
+ ||a|| <= sqrt(1 + SNR ||h||^2).
130
+ """
131
+ h = np.asarray(h, dtype=np.float64)
132
+ snr = 10 ** (snr_db / 10)
133
+ return float(np.sqrt(1.0 + snr * float(h @ h)))
134
+
135
+
136
+ # ─────────────────────────────────────────────────────────────────────────────
137
+ # The fading realization
138
+ # ─────────────────────────────────────────────────────────────────────────────
139
+
140
+ def fading_gains(num_nodes, rng, scale=1.0):
141
+ """
142
+ One real fading realization: h_i ~ N(0, scale^2), i.i.d.
143
+
144
+ Real baseband fading — the complex case (Gaussian-integer lattices)
145
+ is the literature's practical choice (Liu & Ling 2016) and the natural
146
+ extension here.
147
+ """
148
+ return rng.normal(0.0, scale, size=int(num_nodes))
149
+
150
+
151
+ # ─────────────────────────────────────────────────────────────────────────────
152
+ # Selection — the three methods
153
+ # ─────────────────────────────────────────────────────────────────────────────
154
+
155
+ def _rounded_candidates(h, snr_db, iterations=6):
156
+ """
157
+ Nearest-integer heuristic: iterate a <- round(alpha(a) * h).
158
+
159
+ Seeded from sign(h) (so h^T a = ||h||_1 > 0 and alpha > 0) and from each
160
+ signed coordinate axis; keeps the best rate seen. Seeds matter: the
161
+ all-ones seed degenerates on a mixed-sign channel (h^T a <= 0 -> alpha
162
+ = 0), which is exactly the case selection exists to handle.
163
+ """
164
+ h = np.asarray(h, dtype=np.float64)
165
+ n = len(h)
166
+ best = None
167
+ best_rate = 0.0
168
+
169
+ def check(a):
170
+ nonlocal best, best_rate
171
+ a = np.asarray(a, dtype=int)
172
+ if not np.any(a):
173
+ return
174
+ rate = computation_rate(h, a, snr_db)
175
+ if best is None or rate > best_rate:
176
+ best, best_rate = a, rate
177
+
178
+ seeds = [np.sign(h).astype(int)]
179
+ for i in range(n):
180
+ e = np.zeros(n, dtype=int)
181
+ e[i] = int(np.sign(h[i]) or 1)
182
+ seeds.append(e)
183
+
184
+ for seed in seeds:
185
+ a = seed
186
+ check(a)
187
+ for _ in range(iterations):
188
+ alpha = mmse_alpha(h, a, snr_db)
189
+ if alpha <= 0:
190
+ break
191
+ a_new = np.rint(alpha * h).astype(int)
192
+ if not np.any(a_new) or np.array_equal(a_new, a):
193
+ break
194
+ a = a_new
195
+ check(a)
196
+
197
+ if best is None:
198
+ best = np.zeros(n, dtype=int)
199
+ i = int(np.argmax(np.abs(h)))
200
+ best[i] = int(np.sign(h[i]) or 1)
201
+ best_rate = computation_rate(h, best, snr_db)
202
+ return best, best_rate
203
+
204
+
205
+ def _exhaustive(h, snr_db, max_component=5, max_vectors=300_000):
206
+ """Exhaustive search of the integer norm ball ||a|| <= bound."""
207
+ h = np.asarray(h, dtype=np.float64)
208
+ n = len(h)
209
+ bound = norm_bound(h, snr_db)
210
+ radius = min(int(np.floor(bound)), int(max_component))
211
+ while radius >= 1 and (2 * radius + 1) ** n > max_vectors:
212
+ radius -= 1
213
+ best = None
214
+ best_rate = 0.0
215
+ for a in itertools.product(range(-radius, radius + 1), repeat=n):
216
+ a_arr = np.asarray(a, dtype=float)
217
+ if not np.any(a_arr):
218
+ continue
219
+ if float(np.linalg.norm(a_arr)) > bound + 1e-9:
220
+ continue
221
+ rate = computation_rate(h, a_arr, snr_db)
222
+ if best is None or rate > best_rate:
223
+ best, best_rate = np.asarray(a, dtype=int), rate
224
+ if best is None:
225
+ best, best_rate = _rounded_candidates(h, snr_db)
226
+ return best, best_rate
227
+
228
+
229
+ def lll_reduce(basis, delta=0.75):
230
+ """
231
+ Classic LLL (Lenstra-Lenstra-Lovasz) basis reduction on a real basis.
232
+
233
+ basis: (n x n) array whose ROWS span the lattice. Returns
234
+ (reduced_basis, U) with reduced_basis = U @ basis and U unimodular
235
+ integer (det = +/-1). Textbook algorithm with Gram-Schmidt recomputed
236
+ per step — O(n^3) per op, trivial at toy dimensions.
237
+ """
238
+ B = np.array(basis, dtype=np.float64)
239
+ n = len(B)
240
+ U = np.eye(n, dtype=int)
241
+
242
+ def gram_schmidt():
243
+ Bs = np.zeros_like(B)
244
+ mu = np.zeros((n, n))
245
+ for i in range(n):
246
+ Bs[i] = B[i]
247
+ for j in range(i):
248
+ mu[i, j] = (B[i] @ Bs[j]) / (Bs[j] @ Bs[j])
249
+ Bs[i] = Bs[i] - mu[i, j] * Bs[j]
250
+ return Bs, mu
251
+
252
+ Bs, mu = gram_schmidt()
253
+ k = 1
254
+ while k < n:
255
+ for j in range(k - 1, -1, -1):
256
+ q = int(round(mu[k, j]))
257
+ if q:
258
+ B[k] = B[k] - q * B[j]
259
+ U[k] = U[k] - q * U[j]
260
+ Bs, mu = gram_schmidt()
261
+ if (Bs[k] @ Bs[k]) >= (delta - mu[k, k - 1] ** 2) * (Bs[k - 1] @ Bs[k - 1]):
262
+ k += 1
263
+ else:
264
+ B[[k - 1, k]] = B[[k, k - 1]]
265
+ U[[k - 1, k]] = U[[k, k - 1]]
266
+ Bs, mu = gram_schmidt()
267
+ k = max(k - 1, 1)
268
+ return B, U
269
+
270
+
271
+ def _lll(h, snr_db):
272
+ """
273
+ LLL-aided selection: reduce the SLV lattice, then enumerate integer
274
+ coordinates in {-1,0,1} of the REDUCED basis (a bounded, small search —
275
+ the LLL-aided enumeration pattern of Liu & Ling 2016; the polynomial-
276
+ complexity exact algorithm is Sahraei & Gastpar 2014).
277
+
278
+ D(a) = a^T M a with M = (I + SNR h h^T)^{-1} = I - c h h^T,
279
+ c = SNR/(1 + SNR ||h||^2). Factor M = L L^T (Cholesky); then
280
+ D(a) = ||L^T a||^2 — the squared length of the lattice point
281
+ sum_j a_j * row_j(L). Reduce the row basis; reduced rows U_k are
282
+ integer combinations to enumerate over: a = sum_k c_k U_k.
283
+ """
284
+ h = np.asarray(h, dtype=np.float64)
285
+ n = len(h)
286
+ snr = 10 ** (snr_db / 10)
287
+ c = snr / (1.0 + snr * float(h @ h))
288
+ M = np.eye(n) - c * np.outer(h, h)
289
+ L = np.linalg.cholesky(M)
290
+ _, U = lll_reduce(L)
291
+ best = None
292
+ best_rate = 0.0
293
+ for coeffs in itertools.product([-1, 0, 1], repeat=n):
294
+ a = np.asarray(coeffs, dtype=int) @ U
295
+ if not np.any(a):
296
+ continue
297
+ rate = computation_rate(h, a, snr_db)
298
+ if best is None or rate > best_rate:
299
+ best, best_rate = np.array(a, dtype=int), rate
300
+ if best is None:
301
+ best, best_rate = _rounded_candidates(h, snr_db)
302
+ return best, best_rate
303
+
304
+
305
+ def select_coefficients(h, snr_db, method="exhaustive"):
306
+ """
307
+ Choose the integer coefficient vector a maximizing the computation rate.
308
+
309
+ h: real channel gains. snr_db: per-node SNR (P / sigma^2, P = L^2/12).
310
+
311
+ method:
312
+ 'exhaustive' — search the Nazer/Gastpar norm ball (exact, exponential)
313
+ 'lll' — LLL-reduced candidate rows (polynomial-time approximation)
314
+ 'rounded' — nearest-integer heuristic round(alpha * h) (baseline)
315
+
316
+ Returns (a, rate) with a an integer numpy array, never all-zero.
317
+ """
318
+ h = np.asarray(h, dtype=np.float64)
319
+ if method == "exhaustive":
320
+ return _exhaustive(h, snr_db)
321
+ if method == "lll":
322
+ return _lll(h, snr_db)
323
+ if method == "rounded":
324
+ return _rounded_candidates(h, snr_db)
325
+ raise ValueError(
326
+ f"unknown selection method '{method}'; "
327
+ "use 'exhaustive', 'lll', or 'rounded'"
328
+ )