Rename simulation scripts to descriptive names

This commit is contained in:
Ki-Ho Lee
2026-08-25 20:26:08 +09:00
parent 1fc9cad834
commit d9f6f4ac53
10 changed files with 8 additions and 8 deletions
+82
View File
@@ -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)")