diff --git a/README.md b/README.md index 67bedf3..0ff192c 100644 --- a/README.md +++ b/README.md @@ -48,9 +48,10 @@ Run the scripts from inside `code/`. All plots are rendered from | Fig. 2 | Per-user MSE, aware vs blind floor | `revision_sims.py E1` | `floor_validation.csv` | | Fig. 3 | Effective sum rate at the CLIP affinity | `revision_sims.py E7a` | `rate_corrected.csv` | | Fig. 4 | Cosine recovery on real BERT+ViT pairs | `fig_real_merged.py`, then `refine_matched.py` | `bertvit_merged.csv` | -| Fig. 5 | Receiver comparison under Rayleigh fading | `revision_sims_gpu.py E2` | `sic_comparison.csv` | -| Fig. 6 | Value of the measured affinity | `revision_sims.py E7a` | `beta_sweep_corrected.csv` | -| Fig. 7 | Multi-user scaling (joint Wiener) | `revision_sims_gpu.py E7c` | `multiuser_corrected.csv` | +| Fig. 5 | Top-1 retrieval with recovered embeddings | `retrieval_real.py` | `retrieval_real.csv` | +| Fig. 6 | Receiver comparison under Rayleigh fading | `revision_sims_gpu.py E2` | `sic_comparison.csv` | +| Fig. 7 | Value of the measured affinity | `revision_sims.py E7a` | `beta_sweep_corrected.csv` | +| Fig. 8 | Multi-user scaling (joint Wiener) | `revision_sims_gpu.py E7c` | `multiuser_corrected.csv` | Quantities quoted in the text but not plotted come from the same drivers: `revision_sims.py E0` writes `theorem_check.csv` (Theorem 1 diff --git a/code/replot_all.py b/code/replot_all.py index ceb5f33..ad624db 100644 --- a/code/replot_all.py +++ b/code/replot_all.py @@ -179,11 +179,40 @@ def fig_multiuser(): save(fig, "fig_multiuser_corrected") +# ------------------------------------------------ fig_retrieval +def fig_retrieval(): + rows = rows_of("retrieval_real") + snr = col(rows, "snr_db") + fig, ax = plt.subplots() + ax.plot(snr, col(rows, "edma"), "o-", color="C3", label=LBL["edma"]) + ax.plot(snr, col(rows, "hybrid"), "^-", color="C2", + label=LBL["hybrid"]) + ax.plot(snr, col(rows, "todma"), "d-.", color="C4", + label=LBL["todma"]) + ax.plot(snr, col(rows, "oma"), "v:", color="C1", label=LBL["oma"]) + ax.plot(snr, col(rows, "genie"), "-", color="gray", lw=1.0, + label=LBL["genie"]) + ax.axhline(1.0 / 16, color="gray", ls=":", lw=0.8) + ax.annotate("chance $1/16$", xy=(10.5, 1.0 / 16 + 0.015), fontsize=7, + color="gray") + ax.set_xlabel("SNR $\\rho$ [dB]") + ax.set_ylabel("Top-1 retrieval accuracy") + ax.set_xlim(snr[0], snr[-1]); ax.set_ylim(0, 1.42) + ax.set_yticks([0.0, 0.2, 0.4, 0.6, 0.8, 1.0]) + hs, ls = ax.get_legend_handles_labels() + order = [1, 4, 0, 2, 3] # long labels share column one + ax.legend([hs[i] for i in order], [ls[i] for i in order], + loc="upper center", ncol=2, columnspacing=0.7, + handlelength=1.3, handletextpad=0.5) + save(fig, "fig_retrieval") + + if __name__ == "__main__": import sys todo = set(sys.argv[1:]) ALL = {"floor": fig_floor, "rate": fig_rate, "beta": fig_beta_sweep, - "sic": fig_sic, "multi": fig_multiuser} + "sic": fig_sic, "multi": fig_multiuser, + "retr": fig_retrieval} for name, fn in ALL.items(): if not todo or name in todo: fn() diff --git a/code/retrieval_real.py b/code/retrieval_real.py new file mode 100644 index 0000000..3701b1d --- /dev/null +++ b/code/retrieval_real.py @@ -0,0 +1,147 @@ +""" +Task-level validation: nearest-neighbor retrieval with recovered +embeddings on the real BERT/ViT pairs. +=================================================================== +For every recovered embedding, rank the 16 clean database embeddings +of the SAME modality by absolute cosine and score top-1 retrieval of +the transmitted item (chance level 1/16). This measures whether the +recovery preserves semantic identity, the operational question behind +the cosine metric of the manuscript. + +Schemes: EDMA (aware Wiener), hybrid (EDMA + stored refinement gate +from data/refine_gates.npz), OMA (equivalent-bandwidth model), +genie-aided SIC bound, ToDMA-adapted (first 40 fading draws). + +Same channel, energy, and seed conventions as fig_real_merged.py. +Writes data/retrieval_real.csv. Run under WSL (torch, CUDA). +""" +from __future__ import annotations +import csv +import math +import time +import numpy as np +import torch + +from fig_real_merged import (load_pairs, haar_t, aware_batch, todma_prepare, + omp_code, todma_run, SNRS, NFADE, NFADE_TOD, + D, DATA, DEV, SEED) + +torch.manual_seed(SEED) +rng = np.random.default_rng(SEED) + + +def refine_apply_single(P1, z): + """z: (b, D) real torch tensor; P1: (D, D) gate.""" + return D * torch.softmax((z @ P1.T) / math.sqrt(D), dim=1) * z + + +def top1(rec, db, idx): + """rec: (b, D) cfloat; db: (n, D) float; returns (b,) 0/1 hits.""" + sims = (rec @ db.T.to(rec.dtype).conj()).abs() # (b, n) + sims = sims / (rec.norm(dim=1, keepdim=True) + * db.norm(dim=1).unsqueeze(0)) + return (sims.argmax(dim=1) == idx).float().cpu().numpy() + + +def main(): + A, B, betas = load_pairs() + npairs = len(A) + gates = np.load(DATA / "refine_gates.npz") + P1 = torch.tensor(gates["P1"], dtype=torch.float32, device=DEV) + tod = todma_prepare() + codes = [(omp_code(tod[0], A[i], tod[3]), + omp_code(tod[0], B[i], tod[3])) for i in range(npairs)] + At = torch.tensor(A, dtype=torch.float32, device=DEV) + Bt = torch.tensor(B, dtype=torch.float32, device=DEV) + gen = torch.Generator(device=DEV).manual_seed(SEED) + nb = len(SNRS) + sigs = torch.tensor(10 ** (-SNRS / 20.0), dtype=torch.float32, + device=DEV) + keys = ("edma", "hybrid", "oma", "genie", "todma") + acc = {k: np.zeros(nb) for k in keys} + cnt = {k: np.zeros(nb) for k in keys} + t0 = time.time() + for i in range(npairs): + bi = float(betas[i]) + e1, e2 = At[i], Bt[i] + c1c, c2c = codes[i] + for f in range(NFADE): + M = haar_t(2, gen) + M1, M2 = M[0], M[1] + Q = M1.T @ M2 + h = (torch.randn(2, generator=gen, device=DEV) + + 1j * torch.randn(2, generator=gen, device=DEV)) \ + / math.sqrt(2) + n = (torch.randn(D, generator=gen, device=DEV) + + 1j * torch.randn(D, generator=gen, device=DEV)) \ + / math.sqrt(2) + n2 = (torch.randn(D, generator=gen, device=DEV) + + 1j * torch.randn(D, generator=gen, device=DEV)) \ + / math.sqrt(2) + r0 = h[0] * (M1 @ e1).to(torch.cfloat) \ + + h[1] * (M2 @ e2).to(torch.cfloat) + r = r0.unsqueeze(0) + sigs.view(-1, 1) * n.unsqueeze(0) + t1 = (M1.T.to(torch.cfloat) @ r.unsqueeze(-1)).squeeze(-1) / h[0] + t2 = (M2.T.to(torch.cfloat) @ r.unsqueeze(-1)).squeeze(-1) / h[1] + c1 = (h[1] / h[0]).item() + c2 = (h[0] / h[1]).item() + v1 = sigs**2 / h[0].abs()**2 + v2 = sigs**2 / h[1].abs()**2 + g1 = aware_batch(t1, Q, bi, c1, v1) + g2 = aware_batch(t2, Q.T, bi, c2, v2) + acc["edma"] += 0.5 * (top1(g1, At, i) + top1(g2, Bt, i)) + hy1 = refine_apply_single(P1, g1.real.float()).to(torch.cfloat) + hy2 = refine_apply_single(P1, g2.real.float()).to(torch.cfloat) + acc["hybrid"] += 0.5 * (top1(hy1, At, i) + top1(hy2, Bt, i)) + o1 = e1.to(torch.cfloat).unsqueeze(0) \ + + math.sqrt(2) * sigs.view(-1, 1) * n.unsqueeze(0) / h[0] + o2 = e2.to(torch.cfloat).unsqueeze(0) \ + + math.sqrt(2) * sigs.view(-1, 1) * n2.unsqueeze(0) / h[1] + acc["oma"] += 0.5 * (top1(o1, At, i) + top1(o2, Bt, i)) + ge1 = (M1.T.to(torch.cfloat) + @ (r - h[1] * (M2 @ e2).to(torch.cfloat)).unsqueeze(-1) + ).squeeze(-1) / h[0] + ge2 = (M2.T.to(torch.cfloat) + @ (r - h[0] * (M1 @ e1).to(torch.cfloat)).unsqueeze(-1) + ).squeeze(-1) / h[1] + acc["genie"] += 0.5 * (top1(ge1, At, i) + top1(ge2, Bt, i)) + for kk in ("edma", "hybrid", "oma", "genie"): + cnt[kk] += 1 + if f < NFADE_TOD: + hnp = (complex(h[0].item()), complex(h[1].item())) + nslots = [(rng.standard_normal(tod[4]) + + 1j * rng.standard_normal(tod[4])) + / math.sqrt(2) for _ in range(tod[3])] + for k, s in enumerate(SNRS): + sig = 10 ** (-s / 20.0) + recs = todma_run(tod, (c1c, c2c), hnp, sig, nslots) + hit = 0.0 + for j, (rec, db, ii) in enumerate( + ((recs[0], A, i), (recs[1], B, i))): + if rec is None: + continue # failed detection: no hit + sims = np.abs(db @ rec) / ( + np.linalg.norm(db, axis=1) + * np.linalg.norm(rec)) + hit += 0.5 * float(int(np.argmax(sims)) == ii) + acc["todma"][k] += hit + cnt["todma"][k] += 1 + print(f" pair {i+1}/{npairs} done ({time.time()-t0:.0f}s)", + flush=True) + for k in keys: + acc[k] /= np.maximum(cnt[k], 1) + + with open(DATA / "retrieval_real.csv", "w", newline="") as f: + w = csv.writer(f) + w.writerow(["snr_db"] + list(keys)) + for k, s in enumerate(SNRS): + w.writerow([s] + [acc[key][k] for key in keys]) + print(f"[OK] wrote {DATA/'retrieval_real.csv'}") + for k, s in enumerate(SNRS): + print(f" {s:4.1f} dB EDMA {acc['edma'][k]:.3f} " + f"hybrid {acc['hybrid'][k]:.3f} ToDMA {acc['todma'][k]:.3f} " + f" OMA {acc['oma'][k]:.3f} genie {acc['genie'][k]:.3f}") + + +if __name__ == "__main__": + main() diff --git a/data/retrieval_real.csv b/data/retrieval_real.csv new file mode 100644 index 0000000..3f1093b --- /dev/null +++ b/data/retrieval_real.csv @@ -0,0 +1,14 @@ +snr_db,edma,hybrid,oma,genie,todma +0.0,0.1915625,0.2253125,0.1278125,0.190625,0.028125 +2.5,0.278125,0.325,0.1884375,0.2840625,0.03671875 +5.0,0.4065625,0.45375,0.2784375,0.4103125,0.05390625 +7.5,0.5365625,0.5765625,0.3965625,0.54,0.0765625 +10.0,0.6625,0.6940625,0.52375,0.6675,0.1390625 +12.5,0.7646875,0.781875,0.6515625,0.76375,0.23359375 +15.0,0.8296875,0.84875,0.7515625,0.833125,0.371875 +17.5,0.8828125,0.8934375,0.828125,0.89,0.515625 +20.0,0.9178125,0.926875,0.881875,0.925625,0.65859375 +22.5,0.94625,0.9503125,0.9175,0.9565625,0.7703125 +25.0,0.96125,0.9625,0.94625,0.9734375,0.84609375 +27.5,0.970625,0.971875,0.9690625,0.983125,0.8921875 +30.0,0.9796875,0.98,0.9809375,0.99125,0.91640625 diff --git a/fig/fig_retrieval.pdf b/fig/fig_retrieval.pdf new file mode 100644 index 0000000..ca4f5da Binary files /dev/null and b/fig/fig_retrieval.pdf differ