#!/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()