Code and stored results for AI-native multi-user semantic communications (JSAC submission)
This commit is contained in:
Executable
+369
@@ -0,0 +1,369 @@
|
||||
#!/usr/bin/env python3
|
||||
# ============================================================
|
||||
# c21_mnist_flat.py
|
||||
#
|
||||
# MNIST transmission over flat Rayleigh fading with AWGN,
|
||||
# replacing the synthetic-symbol flat study. Compared schemes:
|
||||
# Transformer SE separation, per-user AE, proposed softmax
|
||||
# joint, proposed signed joint, and proposed signed MAML with
|
||||
# decoder-side task-conditional adaptation (warm start from
|
||||
# the signed joint model). The digital chain is
|
||||
# classifier-limited on this channel and is reported from the
|
||||
# transmit-side classifier accuracy.
|
||||
# ============================================================
|
||||
|
||||
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
|
||||
|
||||
import c11_doppler_csi as base
|
||||
import c13_mnist as m13
|
||||
|
||||
|
||||
def apply_channel(y, snr_db):
|
||||
"""Flat real Rayleigh + AWGN with global power normalization."""
|
||||
h = torch.randn(y.size(0), 1, device=y.device) / math.sqrt(2.0)
|
||||
y = h * y
|
||||
y = y / torch.sqrt(torch.mean(y ** 2) + 1e-12)
|
||||
noise_var = 10 ** (-snr_db / 10.0)
|
||||
return y + torch.randn_like(y) * math.sqrt(noise_var)
|
||||
|
||||
|
||||
class FlatBase(nn.Module):
|
||||
"""CNN encoder + masks; rx defined by subclasses."""
|
||||
|
||||
def __init__(self, U, d=128, hidden=256, n_cls=10):
|
||||
super().__init__()
|
||||
self.U, self.d = U, d
|
||||
self.encoder = m13.CNNTrunk(d)
|
||||
self.masks = nn.Parameter(torch.randn(U, d))
|
||||
self.cls = nn.Sequential(
|
||||
nn.Linear(d, hidden), nn.ReLU(inplace=True),
|
||||
nn.Linear(hidden, n_cls))
|
||||
|
||||
def tx(self, imgs):
|
||||
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)
|
||||
return y, m
|
||||
|
||||
def forward(self, imgs, snr_db):
|
||||
y, m = self.tx(imgs)
|
||||
y = apply_channel(y, snr_db)
|
||||
R = y.unsqueeze(1) * m.unsqueeze(0)
|
||||
return self.rx(R)
|
||||
|
||||
|
||||
class FlatSigned(FlatBase):
|
||||
def __init__(self, U, d=128, hidden=256, n_cls=10, score_hidden=64):
|
||||
super().__init__(U, d, hidden, n_cls)
|
||||
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))
|
||||
|
||||
def rx(self, R):
|
||||
B = R.shape[0]
|
||||
T = torch.bmm(R, R.transpose(1, 2)) / self.d
|
||||
w = self.score(T.reshape(B, -1)).view(B, self.U, self.U)
|
||||
W = torch.eye(self.U, device=R.device).unsqueeze(0) + w
|
||||
return self.cls(torch.bmm(W, R))
|
||||
|
||||
|
||||
class FlatSoftmax(FlatBase):
|
||||
def __init__(self, U, d=128, hidden=256, n_cls=10):
|
||||
super().__init__(U, d, hidden, n_cls)
|
||||
self.query = nn.Parameter(torch.randn(U, d))
|
||||
self.key = nn.Linear(d, d, bias=False)
|
||||
self.val = nn.Linear(d, d, bias=False)
|
||||
|
||||
def rx(self, R):
|
||||
K, Vv = self.key(R), self.val(R)
|
||||
q = F.normalize(self.query, dim=1)
|
||||
scores = torch.einsum("ud,bid->bui", q,
|
||||
F.normalize(K, dim=-1)) / math.sqrt(self.d)
|
||||
attn = F.softmax(scores * 1.43, dim=-1)
|
||||
z = torch.einsum("bui,bid->bud", attn, Vv) + R
|
||||
return self.cls(z)
|
||||
|
||||
|
||||
class FlatTransformer(FlatBase):
|
||||
def __init__(self, U, d=128, hidden=256, n_cls=10, 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(d, hidden)
|
||||
self.sep = nn.TransformerEncoder(layer, num_layers=n_layers)
|
||||
self.out = nn.Linear(hidden, n_cls)
|
||||
|
||||
def rx(self, R):
|
||||
return self.out(self.sep(self.inp(R)))
|
||||
|
||||
|
||||
class FlatPerUserAE(nn.Module):
|
||||
"""Per-user CNN encoders and heads, no masks."""
|
||||
|
||||
def __init__(self, U, d=128, hidden=256, n_cls=10):
|
||||
super().__init__()
|
||||
self.U, self.d = U, d
|
||||
self.encs = nn.ModuleList([m13.CNNTrunk(d) for _ in range(U)])
|
||||
self.heads = nn.ModuleList([
|
||||
nn.Sequential(nn.Linear(d, hidden), nn.ReLU(inplace=True),
|
||||
nn.Linear(hidden, n_cls))
|
||||
for _ in range(U)
|
||||
])
|
||||
|
||||
def forward(self, imgs, snr_db):
|
||||
e = torch.stack([self.encs[u](imgs[:, u]) for u in range(self.U)],
|
||||
dim=1)
|
||||
y = e.sum(dim=1)
|
||||
y = apply_channel(y, snr_db)
|
||||
return torch.stack([h(y) for h in self.heads], dim=1)
|
||||
|
||||
|
||||
MODELS = {
|
||||
"transformer": FlatTransformer,
|
||||
"peruser": FlatPerUserAE,
|
||||
"softmax": FlatSoftmax,
|
||||
"signed": FlatSigned,
|
||||
}
|
||||
|
||||
DECODER_PREFIX = ("score.", "cls.")
|
||||
|
||||
|
||||
def train_model(key, args, device, x_tr, y_tr):
|
||||
base.set_seed(args.seed)
|
||||
model = MODELS[key](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))
|
||||
imgs, labels = m13.sample_frames(x_tr, y_tr, args.batch, args.users,
|
||||
device, rng)
|
||||
logits = model(imgs, snr)
|
||||
loss = F.cross_entropy(logits.reshape(-1, 10), labels.reshape(-1))
|
||||
opt.zero_grad(set_to_none=True)
|
||||
loss.backward()
|
||||
opt.step()
|
||||
if step % 2000 == 0:
|
||||
print(f"[{key} {step}/{args.steps}] loss={loss.item():.4f}",
|
||||
flush=True)
|
||||
return model
|
||||
|
||||
|
||||
def rx_with_fast(model, R, fast):
|
||||
B = R.shape[0]
|
||||
T = torch.bmm(R, R.transpose(1, 2)) / model.d
|
||||
h1 = torch.relu(F.linear(T.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"]))
|
||||
w = F.linear(h1, fast["score.4.weight"], fast["score.4.bias"])
|
||||
W = torch.eye(model.U, device=R.device).unsqueeze(0)
|
||||
W = W + w.view(B, model.U, model.U)
|
||||
z = torch.bmm(W, R)
|
||||
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 adapt_decoder(model, x_tr, y_tr, 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(DECODER_PREFIX)}
|
||||
for _ in range(args.eval_inner_steps):
|
||||
imgs, labels = m13.sample_frames(x_tr, y_tr, args.support,
|
||||
args.users, device, rng)
|
||||
with torch.enable_grad():
|
||||
y, m = model.tx(imgs)
|
||||
y = apply_channel(y, snr)
|
||||
R = y.unsqueeze(1) * m.unsqueeze(0)
|
||||
logits = rx_with_fast(model, R, fast)
|
||||
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 meta_train_decoder(model, x_tr, y_tr, args, device):
|
||||
"""First-order MAML on the decoder-side parameters, warm-started
|
||||
from the jointly trained signed model."""
|
||||
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))
|
||||
fast = adapt_decoder(model, x_tr, y_tr, args, snr, device)
|
||||
fast = {k: v.requires_grad_(True) for k, v in fast.items()}
|
||||
imgs, labels = m13.sample_frames(x_tr, y_tr, args.support,
|
||||
args.users, device, rng)
|
||||
ytx, m = model.tx(imgs)
|
||||
ytx = apply_channel(ytx, snr)
|
||||
R = ytx.unsqueeze(1) * m.unsqueeze(0)
|
||||
logits = rx_with_fast(model, R, fast)
|
||||
loss = F.cross_entropy(logits.reshape(-1, 10),
|
||||
labels.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)
|
||||
return model
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def eval_model(model, x, y, args, snr, device, rng, fast=None):
|
||||
errs, total = 0, 0
|
||||
for _ in range(args.eval_nb):
|
||||
imgs, labels = m13.sample_frames(x, y, args.eval_batch, args.users,
|
||||
device, rng)
|
||||
if fast is None:
|
||||
logits = model(imgs, snr)
|
||||
else:
|
||||
ytx, m = model.tx(imgs)
|
||||
ytx = apply_channel(ytx, snr)
|
||||
R = ytx.unsqueeze(1) * m.unsqueeze(0)
|
||||
logits = rx_with_fast(model, R, fast)
|
||||
errs += (logits.argmax(-1) != labels).sum().item()
|
||||
total += labels.numel()
|
||||
return errs / total
|
||||
|
||||
|
||||
def eval_maml_only(args, device):
|
||||
"""Precision re-evaluation of the flat signed MAML receiver."""
|
||||
tr, te = m13.get_datasets(args.data_root)
|
||||
x_tr, y_tr = m13.tensorize(tr)
|
||||
x_te, y_te = m13.tensorize(te)
|
||||
model = FlatSigned(args.users, args.dim, args.hidden).to(device)
|
||||
model.load_state_dict(torch.load(
|
||||
os.path.join(args.save_dir, "flat_maml.pt"), map_location=device))
|
||||
model.eval()
|
||||
joint = FlatSigned(args.users, args.dim, args.hidden).to(device)
|
||||
joint.load_state_dict(torch.load(
|
||||
os.path.join(args.save_dir, "flat_signed.pt"), map_location=device))
|
||||
joint.eval()
|
||||
for snr in args.eval_snrs:
|
||||
sers = []
|
||||
for rep in range(args.adapt_reps):
|
||||
args_seed = args.seed + 77 + rep
|
||||
rng = np.random.default_rng(args_seed)
|
||||
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):
|
||||
imgs, labels = m13.sample_frames(x_tr, y_tr, args.support,
|
||||
args.users, device, rng)
|
||||
with torch.enable_grad():
|
||||
y, m = model.tx(imgs)
|
||||
y = apply_channel(y, float(snr))
|
||||
R = y.unsqueeze(1) * m.unsqueeze(0)
|
||||
logits = rx_with_fast(model, R, fast)
|
||||
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)}
|
||||
fast = {k: v.detach() for k, v in fast.items()}
|
||||
ser = eval_model(model, x_te, y_te, args, float(snr), device,
|
||||
np.random.default_rng(args.seed + 3),
|
||||
fast=fast)
|
||||
sers.append(ser)
|
||||
sj = eval_model(joint, x_te, y_te, args, float(snr), device,
|
||||
np.random.default_rng(args.seed + 3))
|
||||
print(f"snr={snr}: joint={sj:.4e} maml mean={np.mean(sers):.4e} "
|
||||
f"reps={[f'{s:.4e}' for s in sers]}", flush=True)
|
||||
|
||||
|
||||
def main():
|
||||
p = argparse.ArgumentParser()
|
||||
p.add_argument("--mode", choices=["run", "eval-maml"], default="run")
|
||||
p.add_argument("--adapt-reps", type=int, default=5)
|
||||
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("--train-snrs", type=float, nargs="+",
|
||||
default=[0, 5, 10, 15, 20, 25, 30])
|
||||
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("--eval-snrs", type=float, nargs="+",
|
||||
default=[10, 20, 30])
|
||||
p.add_argument("--eval-inner-steps", type=int, default=5)
|
||||
p.add_argument("--eval-inner-lr", type=float, default=0.01)
|
||||
p.add_argument("--meta-steps", type=int, default=1500)
|
||||
p.add_argument("--meta-batch", type=int, default=4)
|
||||
p.add_argument("--support", type=int, default=32)
|
||||
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("--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 == "eval-maml":
|
||||
eval_maml_only(args, device)
|
||||
return
|
||||
|
||||
tr, te = m13.get_datasets(args.data_root)
|
||||
x_tr, y_tr = m13.tensorize(tr)
|
||||
x_te, y_te = m13.tensorize(te)
|
||||
|
||||
csv_path = os.path.join(args.save_dir, "mnist_flat.csv")
|
||||
with open(csv_path, "w", newline="") as f:
|
||||
w = csv.writer(f)
|
||||
w.writerow(["method", "snr_db", "ser"])
|
||||
signed_model = None
|
||||
for key in MODELS:
|
||||
model = train_model(key, args, device, x_tr, y_tr)
|
||||
if key == "signed":
|
||||
signed_model = model
|
||||
torch.save(model.state_dict(),
|
||||
os.path.join(args.save_dir, "flat_signed.pt"))
|
||||
for snr in args.eval_snrs:
|
||||
ser = eval_model(model, x_te, y_te, args, float(snr),
|
||||
device, np.random.default_rng(args.seed + 3))
|
||||
w.writerow([key, snr, ser])
|
||||
f.flush()
|
||||
print(f"[{key}] snr={snr} ser={ser:.4e}", flush=True)
|
||||
# signed MAML: warm-started decoder-side meta-training, then
|
||||
# task-conditional adaptation per evaluation SNR
|
||||
signed_model = meta_train_decoder(signed_model, x_tr, y_tr, args,
|
||||
device)
|
||||
torch.save(signed_model.state_dict(),
|
||||
os.path.join(args.save_dir, "flat_maml.pt"))
|
||||
for snr in args.eval_snrs:
|
||||
fast = adapt_decoder(signed_model, x_tr, y_tr, args,
|
||||
float(snr), device)
|
||||
ser = eval_model(signed_model, x_te, y_te, args, float(snr),
|
||||
device, np.random.default_rng(args.seed + 3),
|
||||
fast=fast)
|
||||
w.writerow(["signed_maml", snr, ser])
|
||||
f.flush()
|
||||
print(f"[signed_maml] snr={snr} ser={ser:.4e}", flush=True)
|
||||
print("saved", csv_path)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user