#!/usr/bin/env python3 # ============================================================ # c20_bert.py # # Real-text semantic transmission with BERT features over the # time-varying TDL channel. Each of the U users transmits one # AG News sentence: a frozen BERT-base encoder produces the # [CLS] feature (768-dim), a shared trainable projection maps # it to the d=128 embedding, and the masked embeddings are # superposed exactly as in the MNIST study. The receiver # recovers each user's news topic (4 classes = 2 bits), and # the conventional digital chain classifies at the # transmitter and sends the 2-bit class index as one QPSK # symbol over the user's comb subcarriers with 16-fold # repetition and MRC. # # Modes: cache -> train / train-tf / train-ae / train-tx-cls # -> train-maml -> eval -> fig # ============================================================ import argparse import csv import os import numpy as np import torch import torch.nn as nn import torch.nn.functional as F import c11_doppler_csi as base import c13_mnist as m13 # ------------------------------------------------------------ # BERT feature cache # ------------------------------------------------------------ def cache_features(args, device): from transformers import AutoTokenizer, AutoModel from datasets import load_dataset tok = AutoTokenizer.from_pretrained("bert-base-uncased") bert = AutoModel.from_pretrained("bert-base-uncased").to(device).eval() ds = load_dataset("fancyzhx/ag_news") out = {} for split, n in [("train", args.cache_train), ("test", args.cache_test)]: texts = ds[split]["text"][:n] labels = torch.tensor(ds[split]["label"][:n]) feats = [] with torch.no_grad(): for i in range(0, len(texts), 128): bt = tok(texts[i:i + 128], padding=True, truncation=True, max_length=64, return_tensors="pt").to(device) cls = bert(**bt).last_hidden_state[:, 0] feats.append(cls.cpu()) if (i // 128) % 20 == 0: print(f"[cache {split}] {i}/{len(texts)}", flush=True) out[split] = (torch.cat(feats), labels) os.makedirs(args.save_dir, exist_ok=True) torch.save(out, os.path.join(args.save_dir, "bert_feats.pt")) print("saved bert_feats.pt", out["train"][0].shape, out["test"][0].shape) def load_features(args): d = torch.load(os.path.join(args.save_dir, "bert_feats.pt"), map_location="cpu") return d["train"], d["test"] def sample_frames(feats, labels, B, U, device, rng): idx = torch.from_numpy(rng.integers(0, feats.shape[0], size=(B * U,))) x = feats[idx].to(device).view(B, U, -1) y = labels[idx].to(device).view(B, U) return x, y # ------------------------------------------------------------ # Models (mirror c13 with a BERT-feature front end) # ------------------------------------------------------------ class BertTrunk(nn.Module): def __init__(self, out_dim, in_dim=768, hidden=256): super().__init__() self.net = nn.Sequential( nn.Linear(in_dim, hidden), nn.ReLU(inplace=True), nn.Linear(hidden, out_dim)) def forward(self, x): return self.net(x) class BertSemanticMA(nn.Module): """Shared projection + masks + signed user-wise attention.""" def __init__(self, U, d=128, hidden=256, n_cls=4, score_hidden=64): super().__init__() self.U, self.d = U, d self.encoder = BertTrunk(d) self.masks = nn.Parameter(torch.randn(U, d)) self.score = nn.Sequential( nn.Linear(U * U, score_hidden), nn.ReLU(inplace=True), nn.Linear(score_hidden, score_hidden), nn.ReLU(inplace=True), nn.Linear(score_hidden, U * U)) self.cls = nn.Sequential( nn.Linear(2 * d, hidden), nn.ReLU(inplace=True), nn.Linear(hidden, n_cls)) def tx(self, x, params=None): B = x.shape[0] e = self.encoder(x.reshape(B * self.U, -1)).view(B, self.U, self.d) m = F.normalize(self.masks, dim=1) y = (e * m.unsqueeze(0)).sum(dim=1) y = y / torch.sqrt(torch.mean(y ** 2, dim=1, keepdim=True) + 1e-12) return y, m def rx(self, Yeq, m, params=None): B = Yeq.shape[0] R = Yeq.unsqueeze(1) * m.unsqueeze(0) phi = torch.cat([R.real, R.imag], dim=-1) T = torch.bmm(phi, phi.transpose(1, 2)) / (2 * self.d) w = self.score(T.reshape(B, -1)).view(B, self.U, self.U) W = torch.eye(self.U, device=phi.device).unsqueeze(0) + w z = torch.bmm(W, phi) return self.cls(z) class BertTransformerMA(BertSemanticMA): def __init__(self, U, d=128, hidden=256, n_cls=4, n_heads=4, n_layers=2): super().__init__(U, d, hidden, n_cls) layer = nn.TransformerEncoderLayer(d_model=hidden, nhead=n_heads, dim_feedforward=2 * hidden, batch_first=True) self.inp = nn.Linear(2 * d, hidden) self.sep = nn.TransformerEncoder(layer, num_layers=n_layers) self.out = nn.Linear(hidden, n_cls) def rx(self, Yeq, m, params=None): R = Yeq.unsqueeze(1) * m.unsqueeze(0) phi = torch.cat([R.real, R.imag], dim=-1) return self.out(self.sep(self.inp(phi))) class BertPerUserAE(nn.Module): """Per-user projections and heads, no masks.""" def __init__(self, U, d=128, hidden=256, n_cls=4): super().__init__() self.U, self.d = U, d self.encs = nn.ModuleList([BertTrunk(d) for _ in range(U)]) self.heads = nn.ModuleList([ nn.Sequential(nn.Linear(2 * d, hidden), nn.ReLU(inplace=True), nn.Linear(hidden, n_cls)) for _ in range(U) ]) def tx(self, x, params=None): e = torch.stack([self.encs[u](x[:, u]) for u in range(self.U)], dim=1) y = e.sum(dim=1) y = y / torch.sqrt(torch.mean(y ** 2, dim=1, keepdim=True) + 1e-12) return y, None def rx(self, Yeq, m, params=None): phi = torch.cat([Yeq.real, Yeq.imag], dim=-1) return torch.stack([h(phi) for h in self.heads], dim=1) class BertTxClassifier(nn.Module): def __init__(self, n_cls=4, hidden=256): super().__init__() self.net = nn.Sequential( nn.Linear(768, hidden), nn.ReLU(inplace=True), nn.Linear(hidden, n_cls)) def forward(self, x): return self.net(x) # ------------------------------------------------------------ # Digital chain: 2-bit class as one QPSK symbol, 16x repetition # ------------------------------------------------------------ def digital_tx(cls_idx, U, N, device): const = torch.tensor([1 + 1j, 1 - 1j, -1 + 1j, -1 - 1j], device=device) / np.sqrt(2.0) B = cls_idx.shape[0] X = torch.zeros(B, N, dtype=torch.complex64, device=device) for u in range(U): ks = torch.arange(u, N, U, device=device) X[:, ks] = const[cls_idx[:, u]].unsqueeze(1) return X def digital_detect(Y, H_hat, U, N): const = torch.tensor([1 + 1j, 1 - 1j, -1 + 1j, -1 - 1j], device=Y.device) / np.sqrt(2.0) B = Y.shape[0] out = torch.zeros(B, U, dtype=torch.long, device=Y.device) for u in range(U): ks = torch.arange(u, N, U, device=Y.device) Z = (torch.conj(H_hat[:, ks]) * Y[:, ks]).sum(dim=1) metric = (Z.unsqueeze(1) * torch.conj(const).unsqueeze(0)).real out[:, u] = metric.argmax(dim=1) return out # ------------------------------------------------------------ # Training / adaptation / evaluation # ------------------------------------------------------------ DECODER_PREFIX = ("score.", "cls.") def train_semantic(args, device, model_cls=BertSemanticMA, ckpt="bert_sem.pt"): (ftr, ltr), _ = load_features(args) chan = base.TDLChannel(args.nfft, args.cp, args.taps, device=device) X_pilot = base.make_pilot(args.nfft, device) base.set_seed(args.seed) model = model_cls(args.users, args.dim, args.hidden).to(device) opt = torch.optim.Adam(model.parameters(), lr=args.lr) rng = np.random.default_rng(args.seed) for step in range(1, args.steps + 1): snr = float(rng.choice(args.train_snrs)) fd = float(rng.choice(args.train_fds)) x, y = sample_frames(ftr, ltr, args.batch, args.users, device, rng) noise_var = 10 ** (-snr / 10.0) g_p, g_d = chan.sample(args.batch, fd, args.delta) logits = m13.forward_frames(model, x, chan, g_p, g_d, X_pilot, noise_var, fd, args.delta) loss = F.cross_entropy(logits.reshape(-1, 4), y.reshape(-1)) opt.zero_grad(set_to_none=True) loss.backward() opt.step() if step % 1000 == 0: print(f"[{ckpt} {step}/{args.steps}] loss={loss.item():.4f}", flush=True) torch.save(model.state_dict(), os.path.join(args.save_dir, ckpt)) print("saved", ckpt) def train_tx_cls(args, device): (ftr, ltr), (fte, lte) = load_features(args) base.set_seed(args.seed) model = BertTxClassifier().to(device) opt = torch.optim.Adam(model.parameters(), lr=1e-3) for ep in range(5): perm = torch.randperm(ftr.shape[0]) for i in range(0, ftr.shape[0], 256): idx = perm[i:i + 256] logits = model(ftr[idx].to(device)) loss = F.cross_entropy(logits, ltr[idx].to(device)) opt.zero_grad(set_to_none=True) loss.backward() opt.step() with torch.no_grad(): acc = 0 for i in range(0, fte.shape[0], 2048): acc += (model(fte[i:i + 2048].to(device)).argmax(-1) == lte[i:i + 2048].to(device)).sum().item() print(f"[tx-cls ep{ep + 1}] test acc={acc / fte.shape[0]:.4f}", flush=True) torch.save(model.state_dict(), os.path.join(args.save_dir, "bert_txcls.pt")) print("saved bert_txcls.pt") def rx_with_fast(model, Yeq, m, fast): B = Yeq.shape[0] R = Yeq.unsqueeze(1) * m.unsqueeze(0) phi = torch.cat([R.real, R.imag], dim=-1) G = torch.bmm(phi, phi.transpose(1, 2)) / (2 * model.d) h1 = torch.relu(F.linear(G.reshape(B, -1), fast["score.0.weight"], fast["score.0.bias"])) h1 = torch.relu(F.linear(h1, fast["score.2.weight"], fast["score.2.bias"])) wsc = F.linear(h1, fast["score.4.weight"], fast["score.4.bias"]) W = torch.eye(model.U, device=Yeq.device).unsqueeze(0) W = W + wsc.view(B, model.U, model.U) z = torch.bmm(W, phi) h2 = torch.relu(F.linear(z, fast["cls.0.weight"], fast["cls.0.bias"])) return F.linear(h2, fast["cls.2.weight"], fast["cls.2.bias"]) def forward_frames_fast(model, fast, x, chan, g_p, g_d, X_pilot, noise_var, fd, delta): y_emb, m = model.tx(x) X_data = y_emb.to(torch.complex64) Y_p = chan.transmit(X_pilot.unsqueeze(0).expand(x.shape[0], -1), g_p, noise_var) Y_d = chan.transmit(X_data, g_d, noise_var) H_ls = Y_p * torch.conj(X_pilot).unsqueeze(0) rho = base.aging_rho(fd, delta, chan.N, chan.cp) H_til = (rho / (1.0 + noise_var)) * H_ls q = 1.0 - (rho ** 2) / (1.0 + noise_var) Yeq = torch.conj(H_til) * Y_d / (H_til.abs() ** 2 + q + noise_var) rms = torch.sqrt(torch.mean(Yeq.abs() ** 2, dim=1, keepdim=True) + 1e-12) return rx_with_fast(model, Yeq / rms, m, fast) def adapt_bert(model, ftr, ltr, chan, X_pilot, args, snr, fd, device): rng = np.random.default_rng(args.seed + 77) fast = {k: v.detach().clone().requires_grad_(True) for k, v in model.named_parameters() if k.startswith(DECODER_PREFIX)} for _ in range(args.eval_inner_steps): x, y = sample_frames(ftr, ltr, args.support, args.users, device, rng) g_p, g_d = chan.sample(args.support, fd, args.delta) noise_var = 10 ** (-snr / 10.0) with torch.enable_grad(): logits = forward_frames_fast(model, fast, x, chan, g_p, g_d, X_pilot, noise_var, fd, args.delta) loss = F.cross_entropy(logits.reshape(-1, 4), y.reshape(-1)) grads = torch.autograd.grad(loss, list(fast.values())) fast = {k: (p - args.eval_inner_lr * g).detach().requires_grad_(True) for (k, p), g in zip(fast.items(), grads)} return {k: v.detach() for k, v in fast.items()} def train_maml(args, device): """Warm-started first-order decoder-side meta-training.""" (ftr, ltr), _ = load_features(args) chan = base.TDLChannel(args.nfft, args.cp, args.taps, device=device) X_pilot = base.make_pilot(args.nfft, device) base.set_seed(args.seed) model = BertSemanticMA(args.users, args.dim, args.hidden).to(device) model.load_state_dict(torch.load(os.path.join(args.save_dir, "bert_sem.pt"), map_location=device)) params = [v for k, v in model.named_parameters() if k.startswith(DECODER_PREFIX)] opt = torch.optim.Adam(params, lr=1e-4) rng = np.random.default_rng(args.seed + 11) for step in range(1, args.meta_steps + 1): opt.zero_grad(set_to_none=True) for _ in range(args.meta_batch): snr = float(rng.choice(args.train_snrs)) fd = float(rng.choice(args.train_fds)) fast = adapt_bert(model, ftr, ltr, chan, X_pilot, args, snr, fd, device) fast = {k: v.requires_grad_(True) for k, v in fast.items()} x, y = sample_frames(ftr, ltr, args.support, args.users, device, rng) g_p, g_d = chan.sample(args.support, fd, args.delta) noise_var = 10 ** (-snr / 10.0) logits = forward_frames_fast(model, fast, x, chan, g_p, g_d, X_pilot, noise_var, fd, args.delta) loss = F.cross_entropy(logits.reshape(-1, 4), y.reshape(-1)) / args.meta_batch grads = torch.autograd.grad(loss, list(fast.values())) named = dict(model.named_parameters()) for (k, _), g in zip(fast.items(), grads): if named[k].grad is None: named[k].grad = g.detach().clone() else: named[k].grad += g.detach() opt.step() if step % 500 == 0: print(f"[meta {step}/{args.meta_steps}]", flush=True) torch.save(model.state_dict(), os.path.join(args.save_dir, "bert_maml.pt")) print("saved bert_maml.pt") @torch.no_grad() def eval_all(args, device): base.set_seed(args.seed + 3) (ftr, ltr), (fte, lte) = load_features(args) chan = base.TDLChannel(args.nfft, args.cp, args.taps, device=device) X_pilot = base.make_pilot(args.nfft, device) def load(cls, name): mdl = cls(args.users, args.dim, args.hidden).to(device) mdl.load_state_dict(torch.load(os.path.join(args.save_dir, name), map_location=device)) mdl.eval() return mdl sem = load(BertSemanticMA, "bert_sem.pt") tfm = load(BertTransformerMA, "bert_tf.pt") aem = load(BertPerUserAE, "bert_ae.pt") mm = load(BertSemanticMA, "bert_maml.pt") txc = BertTxClassifier().to(device) txc.load_state_dict(torch.load(os.path.join(args.save_dir, "bert_txcls.pt"), map_location=device)) txc.eval() rng = np.random.default_rng(args.seed + 3) csv_path = os.path.join(args.save_dir, "bert_results.csv") with open(csv_path, "w", newline="") as f: w = csv.writer(f) w.writerow(["snr_db", "fd_norm", "method", "ser"]) for snr in args.eval_snrs: noise_var = 10 ** (-snr / 10.0) errs = {"digital_genie": 0, "digital_pilot": 0, "semantic": 0, "semantic_tf": 0, "semantic_ae": 0, "semantic_maml": 0} fast_mm = adapt_bert(mm, ftr, ltr, chan, X_pilot, args, snr, args.eval_fd, device) total = 0 for _ in range(args.eval_nb): x, y = sample_frames(fte, lte, args.eval_batch, args.users, device, rng) g_p, g_d = chan.sample(args.eval_batch, args.eval_fd, args.delta) for key, mdl in [("semantic", sem), ("semantic_tf", tfm), ("semantic_ae", aem)]: lg = m13.forward_frames(mdl, x, chan, g_p, g_d, X_pilot, noise_var, args.eval_fd, args.delta) errs[key] += (lg.argmax(-1) != y).sum().item() lg = forward_frames_fast(mm, fast_mm, x, chan, g_p, g_d, X_pilot, noise_var, args.eval_fd, args.delta) errs["semantic_maml"] += (lg.argmax(-1) != y).sum().item() B = x.shape[0] pred = txc(x.reshape(B * args.users, -1)).argmax(-1) pred = pred.view(B, args.users) X_d = digital_tx(pred, args.users, args.nfft, device) Y_d = chan.transmit(X_d, g_d, noise_var) Y_p = chan.transmit(X_pilot.unsqueeze(0).expand(B, -1), g_p, noise_var) H_ls = Y_p * torch.conj(X_pilot).unsqueeze(0) H_true = chan.genie_H(g_d) for name, H in [("digital_genie", H_true), ("digital_pilot", H_ls)]: rec = digital_detect(Y_d, H, args.users, args.nfft) errs[name] += (rec != y).sum().item() total += y.numel() for name, e in errs.items(): w.writerow([snr, args.eval_fd, name, e / total]) f.flush() print(f"snr={snr:5.1f} | " + " ".join(f"{k}={v / total:.4e}" for k, v in errs.items()), flush=True) print("saved", csv_path) def make_fig(args): import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt rows = list(csv.DictReader(open(os.path.join(args.save_dir, "bert_results.csv")))) LAB = {"digital_genie": "Digital chain genie CSI", "digital_pilot": "Digital chain pilot CSI", "semantic_tf": "Transformer SE separation", "semantic_ae": "Per-user AE", "semantic": "Proposed signed joint", "semantic_maml": "Proposed signed MAML"} STY = {"digital_genie": dict(color="gray", marker="^", ls="--"), "digital_pilot": dict(color="k", marker="v", ls="-"), "semantic_tf": dict(color="tab:purple", marker="P", ls="-"), "semantic_ae": dict(color="tab:brown", marker="X", ls="-"), "semantic": dict(color="tab:red", marker="o", ls="-"), "semantic_maml": dict(color="tab:green", marker="D", ls="--")} plt.rcParams.update({"font.size": 13, "axes.labelsize": 13, "xtick.labelsize": 12, "ytick.labelsize": 12, "axes.linewidth": 1.1, "grid.linewidth": 0.8, "xtick.major.width": 1.1, "ytick.major.width": 1.1, "xtick.minor.width": 0.8, "ytick.minor.width": 0.8, "xtick.major.size": 4.5, "ytick.major.size": 4.5}) fig = plt.figure(figsize=(5.2, 3.9)) ax = fig.add_axes([0.185, 0.145, 0.79, 0.79]) for mkey in LAB: pts = sorted([(float(r["snr_db"]), float(r["ser"])) for r in rows if r["method"] == mkey]) pts = [(a, b) for a, b in pts if b > 0] if not pts: continue xs, ys = zip(*pts) ax.semilogy(xs, ys, label=LAB[mkey], ms=5, lw=1.8, **STY[mkey]) ax.set_xlabel("SNR (dB)") ax.set_ylabel("SER") # the SER here spans less than one decade, so the default log labels # are the wide "6 x 10^-1" form, which pushes the y label off canvas; # plain decimals keep the axis narrow and the label inside from matplotlib.ticker import FixedLocator, FixedFormatter, NullFormatter yt = [0.1, 0.15, 0.2, 0.3, 0.4, 0.6] ax.yaxis.set_major_locator(FixedLocator(yt)) ax.yaxis.set_major_formatter(FixedFormatter([f"{v:g}" for v in yt])) ax.yaxis.set_minor_formatter(NullFormatter()) ax.grid(True, which="both", alpha=0.35) ax.legend(fontsize=9, framealpha=1.0, labelspacing=0.3, handlelength=1.8, loc="center right", bbox_to_anchor=(0.985, 0.66)) out = os.path.join(args.fig_dir, f"bert_ser_vs_snr_fd{args.eval_fd}.pdf") fig.savefig(out) print("saved", out) def main(): p = argparse.ArgumentParser() p.add_argument("--mode", choices=["cache", "train", "train-tf", "train-ae", "train-tx-cls", "train-maml", "eval", "fig"], required=True) p.add_argument("--users", type=int, default=8) p.add_argument("--dim", type=int, default=128) p.add_argument("--hidden", type=int, default=256) p.add_argument("--nfft", type=int, default=128) p.add_argument("--cp", type=int, default=16) p.add_argument("--taps", type=int, default=8) p.add_argument("--delta", type=int, default=6) p.add_argument("--train-snrs", type=float, nargs="+", default=[0, 5, 10, 15, 20, 25]) p.add_argument("--train-fds", type=float, nargs="+", default=[0.002, 0.005, 0.01, 0.02, 0.05, 0.1]) p.add_argument("--steps", type=int, default=10000) p.add_argument("--batch", type=int, default=64) p.add_argument("--lr", type=float, default=1e-3) p.add_argument("--meta-steps", type=int, default=1500) p.add_argument("--meta-batch", type=int, default=4) p.add_argument("--eval-inner-steps", type=int, default=5) p.add_argument("--eval-inner-lr", type=float, default=0.01) p.add_argument("--support", type=int, default=32) p.add_argument("--eval-snrs", type=float, nargs="+", default=[0, 5, 10, 15, 20, 25, 30]) p.add_argument("--eval-fd", type=float, default=0.05) p.add_argument("--eval-batch", type=int, default=64) p.add_argument("--eval-nb", type=int, default=60) p.add_argument("--cache-train", type=int, default=20000) p.add_argument("--cache-test", type=int, default=7600) p.add_argument("--seed", type=int, default=0) p.add_argument("--save-dir", type=str, default="results_bert") p.add_argument("--fig-dir", type=str, default="fig") args = p.parse_args() device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print("device:", device) if args.mode == "cache": cache_features(args, device) elif args.mode == "train": train_semantic(args, device) elif args.mode == "train-tf": train_semantic(args, device, model_cls=BertTransformerMA, ckpt="bert_tf.pt") elif args.mode == "train-ae": train_semantic(args, device, model_cls=BertPerUserAE, ckpt="bert_ae.pt") elif args.mode == "train-tx-cls": train_tx_cls(args, device) elif args.mode == "train-maml": train_maml(args, device) elif args.mode == "eval": eval_all(args, device) elif args.mode == "fig": make_fig(args) if __name__ == "__main__": main()