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.
541 lines
23 KiB
Python
Executable File
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()
|