Files
JSAC_AIRAN/c20_bert.py
T
KiHoLee 6c8471ece0 Match figure styling and labels to the submitted manuscript
Enlarge the in-canvas fonts and line weights of the result figures so
that they stay legible at the printed column width, split the Fig. 5
convergence curve into a pre-meta adaptation entry and the proposed
MAML entry, and rename the autoencoder legend to match the table row.
Add the analytic complexity replot behind Fig. 4, which was missing
from the repository, and correct the table numbering in the README.
2026-08-03 12:49:53 +09:00

541 lines
23 KiB
Python
Executable File

#!/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.155, 0.145, 0.82, 0.82])
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")
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()