"""Real-data (digits) training: UWCA decoder WITH and WITHOUT MAML, plus OFDMA/ SFDMA and NOMA-SIC, measuring downstream classification accuracy vs SNR for HIGH/LOW/MIX. Produces results/fig_realdata_c.pdf with 4 curves.""" import types, numpy as np, torch import matplotlib; matplotlib.use('Agg'); import matplotlib.pyplot as plt import maml_semantic as M from sklearn.datasets import load_digits dev='cpu'; U,D=4,64; DPU=D//U MASKS=np.zeros((U,D)) for u in range(U): MASKS[u,u*DPU:(u+1)*DPU]=1.0 NOMA_POWER=np.array([0.40,0.30,0.20,0.10]) X,y=load_digits(return_X_y=True); X=X.astype(np.float64); X=X-X.mean(0,keepdims=True) X=X/(np.linalg.norm(X,axis=1,keepdims=True)+1e-8) by_class={c:X[y==c] for c in range(10)} PROTO=np.stack([by_class[c].mean(0) for c in range(10)]); PROTO=PROTO/(np.linalg.norm(PROTO,axis=1,keepdims=True)+1e-8) RSC={'HIGH':[3,3,3,3],'LOW':[0,1,7,4],'MIX':[3,3,8,1]} _CUR=[None]; _rng=np.random.default_rng(0) def my_gen(n,d=64,U=4,rng=None,scenario_cfg=None): ca=_CUR[0] out=np.stack([by_class[ca[u]][_rng.integers(0,len(by_class[ca[u]]),n)] for u in range(U)],1) return torch.tensor(out,dtype=torch.float32) M.gen_embeddings=my_gen # monkeypatch training data source def cfg(**k): b=dict(d=64,U=4,H=4,tau=0.45,lam=0.1,snr_min=0.0,snr_max=20.0,snr_step=2.0, inner_lr=0.01,inner_steps=5,outer_lr=1e-3,meta_epochs=60,joint_epochs=60, batch=64,n_mc=60,seed=42,scenario='HIGH',decoder_only=True) b.update(k); return types.SimpleNamespace(**b) def _norm(E): return E/(np.linalg.norm(E,axis=-1,keepdims=True)+1e-8) def acc_np(Eh, ca): return (np.einsum('nud,cd->nuc',_norm(Eh),PROTO).argmax(-1)==np.array(ca)[None,:]).mean() def se_np(E,snr): n=E.shape[0]; Ytx=(E*MASKS[None]).sum(1) h=np.sqrt(_rng.standard_normal((n,U,1))**2+_rng.standard_normal((n,U,1))**2)*np.sqrt(0.5) return h*Ytx[:,None,:]+_rng.standard_normal((n,U,D))*np.sqrt(float(np.mean(Ytx**2))/(10**(snr/10))) def ofdma_np(Y): return np.stack([_norm(Y[:,u,:]*MASKS[u]) for u in range(U)],1) def noma_np(E,snr): n=E.shape[0]; h=np.sqrt(_rng.standard_normal((n,U,1))**2+_rng.standard_normal((n,U,1))**2)*np.sqrt(0.5) yv=(E*np.sqrt(NOMA_POWER)[None,:,None]*h).sum(1); yv=yv+_rng.standard_normal((n,D))*np.sqrt(float(np.mean(yv**2))/(10**(snr/10))) Eh=np.zeros((n,U,D)); r=yv.copy() for u in range(U): Eh[:,u,:]=_norm(r/(h[:,u,:]+1e-8)); r-=h[:,u,:]*np.sqrt(NOMA_POWER[u])*Eh[:,u,:] return Eh def samp(n,ca): return np.stack([by_class[ca[u]][_rng.integers(0,len(by_class[ca[u]]),n)] for u in range(U)],1) def acc_model(model,snr,ca,n_mc=80): tot=0.0 for _ in range(n_mc): Xb=torch.tensor(samp(64,ca),dtype=torch.float32).to(dev) with torch.no_grad(): E,Ehat,_=model(Xb,float(snr)) tot+=acc_np(Ehat.cpu().numpy(),ca) return tot/n_mc SNR=np.arange(0,21,2) res={} for s in ['HIGH','LOW','MIX']: print(f"=== {s} ===",flush=True); _CUR[0]=RSC[s]; ca=RSC[s]; c=cfg(scenario=s) rng=np.random.default_rng(42) mm=M.SemanticCommSystem(c.d,c.U,c.H,decoder_only=True).to(dev); M.MAMLTrainer(mm,c,dev,rng,None).train() mj=M.SemanticCommSystem(c.d,c.U,c.H,decoder_only=True).to(dev); M.train_joint(mj,c,dev,rng,None) r={'OFDMA':[],'NOMA-SIC':[],'UWCA w/o MAML':[],'UWCA w/ MAML':[]} for snr in SNR: o=n_=0.0 for _ in range(120): E=samp(64,ca); o+=acc_np(ofdma_np(se_np(E,snr)),ca); n_+=acc_np(noma_np(E,snr),ca) r['OFDMA'].append(o/120); r['NOMA-SIC'].append(n_/120) r['UWCA w/o MAML'].append(acc_model(mj,snr,ca)); r['UWCA w/ MAML'].append(acc_model(mm,snr,ca)) res[s]=r print(f" {s} @20dB: OFDMA={r['OFDMA'][-1]:.3f} NOMA={r['NOMA-SIC'][-1]:.3f} w/oMAML={r['UWCA w/o MAML'][-1]:.3f} w/MAML={r['UWCA w/ MAML'][-1]:.3f}",flush=True) import json; json.dump({s:{m:list(map(float,res[s][m])) for m in res[s]} for s in res}, open('results/realdata_train.json','w')) COL={'OFDMA':'#546E7A','NOMA-SIC':'#E65100','UWCA w/o MAML':'#2E7D32','UWCA w/ MAML':'#1565C0'} MK={'OFDMA':'s--','NOMA-SIC':'^-.','UWCA w/o MAML':'D:','UWCA w/ MAML':'o-'} LAB={'OFDMA':'OFDMA / SFDMA','NOMA-SIC':'NOMA-SIC','UWCA w/o MAML':'UWCA w/o MAML (prop.)','UWCA w/ MAML':'UWCA w/ MAML (prop.)'} fig,ax=plt.subplots(1,3,figsize=(11,3.2)) for j,s in enumerate(['HIGH','LOW','MIX']): for m in COL: ax[j].plot(SNR,res[s][m],MK[m],color=COL[m],lw=2,ms=5,label=LAB[m]) ax[j].text(0.5,-0.34,f"({chr(97+j)}) {s}",transform=ax[j].transAxes,ha='center',fontsize=11) ax[j].set_ylim(0.1,1.0); ax[j].set_xlim(0,20); ax[j].grid(alpha=.3); ax[j].set_box_aspect(0.8); ax[j].set_xlabel('SNR (dB)',fontsize=10) if j==0: ax[j].set_ylabel('Downstream accuracy',fontsize=10); ax[j].legend(fontsize=7.5,loc='lower right') fig.tight_layout(); fig.savefig('results/fig_realdata_c.pdf',bbox_inches='tight'); print("saved fig_realdata_c.pdf (trained w/ and w/o MAML)")