Code and stored results for AI-native multi-user semantic communications (JSAC submission)

This commit is contained in:
KiHoLee
2026-08-02 16:49:49 +09:00
commit bed475c954
31 changed files with 3202 additions and 0 deletions
+369
View File
@@ -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()