Files
JSAC_AIRAN/c13_mnist.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

594 lines
25 KiB
Python
Executable File

#!/usr/bin/env python3
# ============================================================
# c13_mnist.py
#
# Real-data validation: multi-user semantic transmission of MNIST
# images over the time-varying frequency-selective channel of c11.
#
# Proposed: shared CNN semantic encoder -> d-dim embedding per user
# -> learnable user masks + over-the-air superposition -> OFDM
# frame (pilot + data, CSI aging) -> aging-aware LMMSE + RMS norm
# -> signed user-wise attention demux -> shared 10-class head.
# Metric: semantic error rate, P(recovered class != true label).
#
# Conventional scheme: transmit-side classification with the same
# CNN trunk, then digital transmission of the 4-bit class index as
# two QPSK symbols over the user's 16 comb subcarriers (8-fold
# repetition each, MRC with genie or aged pilot CSI).
#
# Usage:
# python3 c13_mnist.py --mode train
# python3 c13_mnist.py --mode train-tx-cls
# python3 c13_mnist.py --mode eval
# python3 c13_mnist.py --mode fig
# ============================================================
import argparse
import csv
import math
import os
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from torchvision import datasets, transforms
import c11_doppler_csi as base
# ------------------------------------------------------------
# Models
# ------------------------------------------------------------
class CNNTrunk(nn.Module):
def __init__(self, out_dim):
super().__init__()
self.net = nn.Sequential(
nn.Conv2d(1, 16, 3, stride=2, padding=1), nn.ReLU(inplace=True),
nn.Conv2d(16, 32, 3, stride=2, padding=1), nn.ReLU(inplace=True),
nn.Flatten(),
nn.Linear(32 * 7 * 7, out_dim),
)
def forward(self, x):
return self.net(x)
class MnistSemanticMA(nn.Module):
"""CNN encoder + masks + signed user-wise attention + classifier."""
def __init__(self, U, d=128, hidden=256, n_cls=10, score_hidden=64):
super().__init__()
self.U, self.d = U, d
self.encoder = CNNTrunk(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, imgs, params=None):
# imgs: (B,U,1,28,28)
B = imgs.shape[0]
e = self.encoder(imgs.reshape(B * self.U, 1, 28, 28)).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) # (B,U,10)
class MnistTransformerMA(MnistSemanticMA):
"""SOTA variant: Transformer separation on the equalized features."""
def __init__(self, U, d=128, hidden=256, n_cls=10, n_heads=4, n_layers=2):
super().__init__(U, d, hidden, n_cls)
import torch.nn as nn
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):
import torch
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 MnistPerUserAE(nn.Module):
"""DeepMA-style per-user AE: dedicated per-user CNN encoders (no
shared masks) and dedicated per-user decoder heads, superposed on
the identical full-band resource with unit transmit power."""
def __init__(self, U, d=128, hidden=256, n_cls=10):
super().__init__()
self.U, self.d = U, d
self.encs = nn.ModuleList([CNNTrunk(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, imgs, params=None):
e = torch.stack([self.encs[u](imgs[:, u]) for u in range(self.U)],
dim=1) # (B,U,d)
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) # (B,2d)
return torch.stack([h(phi) for h in self.heads], dim=1)
class TxClassifier(nn.Module):
"""Transmit-side classifier for the conventional digital chain."""
def __init__(self, d=128, n_cls=10):
super().__init__()
self.trunk = CNNTrunk(d)
self.head = nn.Linear(d, n_cls)
def forward(self, x):
return self.head(F.relu(self.trunk(x)))
# ------------------------------------------------------------
# Data
# ------------------------------------------------------------
def get_datasets(root):
tf = transforms.Compose([transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,))])
tr = datasets.MNIST(root, train=True, download=True, transform=tf)
te = datasets.MNIST(root, train=False, download=True, transform=tf)
return tr, te
def sample_frames(ds_x, ds_y, B, U, device, rng):
idx = torch.from_numpy(rng.integers(0, ds_x.shape[0], size=(B * U,)))
imgs = ds_x[idx].to(device).view(B, U, 1, 28, 28)
labels = ds_y[idx].to(device).view(B, U)
return imgs, labels
def tensorize(ds):
x = torch.stack([img for img, _ in ds])
y = torch.tensor([lbl for _, lbl in ds])
return x, y
# ------------------------------------------------------------
# End-to-end forward (mirrors base.semantic_forward)
# ------------------------------------------------------------
def forward_frames(model, imgs, chan, g_p, g_d, X_pilot, noise_var, fd, delta):
y_emb, m = model.tx(imgs)
X_data = y_emb.to(torch.complex64)
Y_p = chan.transmit(X_pilot.unsqueeze(0).expand(imgs.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 model.rx(Yeq / rms, m)
# ------------------------------------------------------------
# Conventional digital chain (TX classification + QPSK index)
# ------------------------------------------------------------
def digital_tx(pred_cls, U, N, device):
"""4-bit class index -> two QPSK symbols on comb subcarriers."""
b32 = (pred_cls >> 2) & 3 # (B,U) first 2 bits
b10 = pred_cls & 3
const = base.qpsk_constellation(device)
s1, s2 = const[b32], const[b10]
B = pred_cls.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) # 16 comb tones
X[:, ks[:8]] = s1[:, u].unsqueeze(1)
X[:, ks[8:]] = s2[:, u].unsqueeze(1)
return X
def digital_detect(Y, H_hat, U, N):
const = base.qpsk_constellation(Y.device)
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)
idx_pair = []
for grp in (ks[:8], ks[8:]):
Z = (torch.conj(H_hat[:, grp]) * Y[:, grp]).sum(dim=1)
metric = torch.real(torch.conj(const).view(1, 4) * Z.unsqueeze(1))
idx_pair.append(metric.argmax(dim=1))
out[:, u] = idx_pair[0] * 4 + idx_pair[1]
return out # (B,U) in 0..15
# ------------------------------------------------------------
# Training / evaluation
# ------------------------------------------------------------
def train_semantic(args, device, model_cls=None, ckpt="mnist_semantic.pt"):
base.set_seed(args.seed)
tr, _ = get_datasets(args.data_root)
x, y = tensorize(tr)
chan = base.TDLChannel(args.nfft, args.cp, args.taps, device=device)
X_pilot = base.make_pilot(args.nfft, device)
if model_cls is None:
model_cls = MnistSemanticMA
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, fd = base.sample_task(args, rng)
imgs, labels = sample_frames(x, y, args.batch, args.users, device, rng)
g_p, g_d = chan.sample(args.batch, fd, args.delta)
noise_var = 10 ** (-snr / 10.0)
logits = forward_frames(model, imgs, chan, g_p, g_d, X_pilot,
noise_var, fd, args.delta)
loss = F.cross_entropy(logits.reshape(-1, 10), labels.reshape(-1))
opt.zero_grad(set_to_none=True)
loss.backward()
opt.step()
if step % 200 == 0:
print(f"[mnist-sem {step}/{args.steps}] loss={loss.item():.4f}", flush=True)
os.makedirs(args.save_dir, exist_ok=True)
torch.save(model.state_dict(), os.path.join(args.save_dir, ckpt))
print("saved " + ckpt)
def train_tx_cls(args, device):
base.set_seed(args.seed)
tr, te = get_datasets(args.data_root)
x, y = tensorize(tr)
xt, yt = tensorize(te)
model = TxClassifier(args.dim).to(device)
opt = torch.optim.Adam(model.parameters(), lr=1e-3)
rng = np.random.default_rng(args.seed)
n = x.shape[0]
for ep in range(3):
perm = torch.from_numpy(rng.permutation(n))
for i in range(0, n, 512):
idx = perm[i:i + 512]
logits = model(x[idx].to(device))
loss = F.cross_entropy(logits, y[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, xt.shape[0], 2048):
acc += (model(xt[i:i + 2048].to(device)).argmax(-1)
== yt[i:i + 2048].to(device)).sum().item()
print(f"[tx-cls epoch {ep + 1}] test acc={acc / xt.shape[0]:.4f}", flush=True)
os.makedirs(args.save_dir, exist_ok=True)
torch.save(model.state_dict(), os.path.join(args.save_dir, "mnist_txcls.pt"))
print("saved mnist_txcls.pt")
MNIST_DECODER_KEYS_PREFIX = ("score.", "cls.")
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, imgs, chan, g_p, g_d, X_pilot,
noise_var, fd, delta):
y_emb, m = model.tx(imgs)
X_data = y_emb.to(torch.complex64)
Y_p = chan.transmit(X_pilot.unsqueeze(0).expand(imgs.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_mnist(model, xs, ys, chan, X_pilot, args, snr, 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(MNIST_DECODER_KEYS_PREFIX)}
for _ in range(args.eval_inner_steps):
imgs, labels = sample_frames(xs, ys, args.support, args.users,
device, rng)
g_p, g_d = chan.sample(args.support, args.eval_fd, args.delta)
noise_var = 10 ** (-snr / 10.0)
with torch.enable_grad():
logits = forward_frames_fast(model, fast, imgs, chan, g_p, g_d,
X_pilot, noise_var, args.eval_fd,
args.delta)
loss = F.cross_entropy(logits.reshape(-1, 10),
labels.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_semantic(args, device):
base.set_seed(args.seed)
tr, _ = get_datasets(args.data_root)
x, y = tensorize(tr)
chan = base.TDLChannel(args.nfft, args.cp, args.taps, device=device)
X_pilot = base.make_pilot(args.nfft, device)
model = MnistSemanticMA(args.users, args.dim, args.hidden).to(device)
warm = os.path.join(args.save_dir, "mnist_semantic.pt")
if os.path.exists(warm):
model.load_state_dict(torch.load(warm, map_location=device))
print("meta-training warm-started from mnist_semantic.pt", flush=True)
opt = torch.optim.Adam(model.parameters(), lr=args.lr)
rng = np.random.default_rng(args.seed)
def task_loss(fast, snr, fd):
imgs, labels = sample_frames(x, y, args.batch, args.users, device,
rng)
g_p, g_d = chan.sample(args.batch, fd, args.delta)
noise_var = 10 ** (-snr / 10.0)
logits = forward_frames_fast(model, fast, imgs, chan, g_p, g_d,
X_pilot, noise_var, fd, args.delta)
return F.cross_entropy(logits.reshape(-1, 10), labels.reshape(-1))
for step in range(1, args.meta_steps + 1):
meta_loss = 0.0
for _ in range(args.meta_batch):
snr, fd = base.sample_task(args, rng)
fast = {k: v for k, v in model.named_parameters()
if k.startswith(MNIST_DECODER_KEYS_PREFIX)}
loss_sup = task_loss(fast, snr, fd)
grads = torch.autograd.grad(loss_sup, list(fast.values()))
fast = {k: p - args.inner_lr * g.detach()
for (k, p), g in zip(fast.items(), grads)}
meta_loss = meta_loss + task_loss(fast, snr, fd)
meta_loss = meta_loss / args.meta_batch
opt.zero_grad(set_to_none=True)
meta_loss.backward()
opt.step()
if step % 100 == 0:
print("[mnist-maml %d/%d] loss=%.4f"
% (step, args.meta_steps, meta_loss.item()), flush=True)
os.makedirs(args.save_dir, exist_ok=True)
torch.save(model.state_dict(),
os.path.join(args.save_dir, "mnist_maml.pt"))
print("saved mnist_maml.pt")
@torch.no_grad()
def eval_all(args, device):
base.set_seed(args.seed + 3)
_, te = get_datasets(args.data_root)
x, y = tensorize(te)
chan = base.TDLChannel(args.nfft, args.cp, args.taps, device=device)
X_pilot = base.make_pilot(args.nfft, device)
sem = MnistSemanticMA(args.users, args.dim, args.hidden).to(device)
sem.load_state_dict(torch.load(os.path.join(args.save_dir, "mnist_semantic.pt"),
map_location=device))
sem.eval()
maml_path = os.path.join(args.save_dir, "mnist_maml.pt")
mm = None
if os.path.exists(maml_path):
mm = MnistSemanticMA(args.users, args.dim, args.hidden).to(device)
mm.load_state_dict(torch.load(maml_path, map_location=device))
mm.eval()
xs_tr, ys_tr = tensorize(get_datasets(args.data_root)[0])
tf_path = os.path.join(args.save_dir, "mnist_tf.pt")
tfm = None
if os.path.exists(tf_path):
tfm = MnistTransformerMA(args.users, args.dim, args.hidden).to(device)
tfm.load_state_dict(torch.load(tf_path, map_location=device))
tfm.eval()
ae_path = os.path.join(args.save_dir, "mnist_ae.pt")
aem = None
if os.path.exists(ae_path):
aem = MnistPerUserAE(args.users, args.dim, args.hidden).to(device)
aem.load_state_dict(torch.load(ae_path, map_location=device))
aem.eval()
txc = TxClassifier(args.dim).to(device)
txc.load_state_dict(torch.load(os.path.join(args.save_dir, "mnist_txcls.pt"),
map_location=device))
txc.eval()
rng = np.random.default_rng(args.seed + 3)
csv_path = os.path.join(args.save_dir, "mnist_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}
if tfm is not None:
errs["semantic_tf"] = 0
if aem is not None:
errs["semantic_ae"] = 0
if mm is not None:
errs["semantic_maml"] = 0
fast_mm = adapt_mnist(mm, xs_tr, ys_tr, chan, X_pilot,
args, snr, device)
total = 0
for _ in range(args.eval_nb):
imgs, labels = sample_frames(x, y, args.eval_batch, args.users,
device, rng)
g_p, g_d = chan.sample(args.eval_batch, args.eval_fd, args.delta)
# proposed semantic chain
logits = forward_frames(sem, imgs, chan, g_p, g_d, X_pilot,
noise_var, args.eval_fd, args.delta)
errs["semantic"] += (logits.argmax(-1) != labels).sum().item()
if tfm is not None:
lg2 = forward_frames(tfm, imgs, chan, g_p, g_d, X_pilot,
noise_var, args.eval_fd, args.delta)
errs["semantic_tf"] += (lg2.argmax(-1) != labels).sum().item()
if aem is not None:
lg4 = forward_frames(aem, imgs, chan, g_p, g_d, X_pilot,
noise_var, args.eval_fd, args.delta)
errs["semantic_ae"] += (lg4.argmax(-1) != labels).sum().item()
if mm is not None:
lg3 = forward_frames_fast(mm, fast_mm, imgs, chan, g_p,
g_d, X_pilot, noise_var,
args.eval_fd, args.delta)
errs["semantic_maml"] += (lg3.argmax(-1) != labels).sum().item()
# conventional digital chain
B = imgs.shape[0]
pred_cls = txc(imgs.reshape(B * args.users, 1, 28, 28)).argmax(-1)
pred_cls = pred_cls.view(B, args.users)
X_d = digital_tx(pred_cls, 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 != labels).sum().item()
total += labels.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(f"saved {csv_path}")
def make_fig(args):
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
rows = []
with open(os.path.join(args.save_dir, "mnist_results.csv")) as f:
rows = list(csv.DictReader(f))
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])
if not pts:
continue
xs, ys = zip(*pts)
ys = [max(v, 1e-5) for v in ys]
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"mnist_ser_vs_snr_fd{args.eval_fd}.pdf")
fig.savefig(out)
print("saved", out)
def main():
p = argparse.ArgumentParser()
p.add_argument("--mode", choices=["train", "train-tf", "train-ae", "train-maml", "train-tx-cls", "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=4000)
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=2500)
p.add_argument("--meta-batch", type=int, default=4)
p.add_argument("--inner-lr", type=float, default=0.02)
p.add_argument("--support", type=int, default=32)
p.add_argument("--eval-inner-steps", type=int, default=5)
p.add_argument("--eval-inner-lr", type=float, default=0.01)
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("--seed", type=int, default=0)
p.add_argument("--save-dir", type=str, default="results_mnist")
p.add_argument("--fig-dir", type=str, default="fig")
p.add_argument("--data-root", type=str, default="data_mnist")
args = p.parse_args()
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print("device:", device)
if args.mode == "train":
train_semantic(args, device)
elif args.mode == "train-maml":
train_maml_semantic(args, device)
elif args.mode == "train-tf":
train_semantic(args, device, model_cls=MnistTransformerMA, ckpt="mnist_tf.pt")
elif args.mode == "train-ae":
train_semantic(args, device, model_cls=MnistPerUserAE, ckpt="mnist_ae.pt")
elif args.mode == "train-tx-cls":
train_tx_cls(args, device)
elif args.mode == "eval":
eval_all(args, device)
elif args.mode == "fig":
make_fig(args)
if __name__ == "__main__":
main()