534 lines
22 KiB
Python
Executable File
534 lines
22 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 multiple access",
|
|
"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="--")}
|
|
fig = plt.figure(figsize=(5.2, 3.9))
|
|
ax = fig.add_axes([0.14, 0.125, 0.835, 0.845])
|
|
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=4, lw=1.3, **STY[mkey])
|
|
ax.set_xlabel("SNR (dB)")
|
|
ax.set_ylabel("SER")
|
|
ax.grid(True, which="both", alpha=0.35)
|
|
ax.legend(fontsize=7.5, 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()
|