Files
JSAC_AIRAN/c20_bert.py
T
KiHoLee 8fe5f499b8 Widen the axes margin and fix the clipped y label of Fig. 8
The BERT study spans less than one decade, so its log axis carried the
wide "6 x 10^-1" tick labels, which pushed the y label off the canvas.
Give that axis plain decimal ticks and widen the axes margin of every
result figure by the same amount, which keeps the axes box identical
across figures at the 4:3 ratio.
2026-08-26 19:09:49 +09:00

549 lines
24 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.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()