Restructure package: descriptive study documentation and clean layout
This commit is contained in:
Executable
+82
@@ -0,0 +1,82 @@
|
||||
"""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)")
|
||||
Reference in New Issue
Block a user