214 lines
6.7 KiB
Matlab
214 lines
6.7 KiB
Matlab
function minimal_model_gpt
|
||
% Minimal binary IM/DD + ISI (M=3), AWGN. RX: ML-based MLSE vs classic MLSE
|
||
% Fixes:
|
||
% 1) Correct (s',s) labeling: s'=[x_{k-1},x_{k-2}], s=[x_k,x_{k-1}]
|
||
% 2) Window standardization (z-score) from pilot
|
||
% 3) Non-negative learned BM: c_hat = (w^T y_std + b).^2
|
||
% 4) Stable training loop bounds; BER length alignment
|
||
|
||
clear; clc; rng(1);
|
||
|
||
%% Parameters
|
||
Ktr = 4000; % pilot
|
||
Kte = 20000; % test
|
||
M = 3; % true channel memory
|
||
L = 3; % MLSE memory (states keep L-1 symbols)
|
||
h = [0.55 1.00 0.45]; % ISI FIR
|
||
EsN0dB= 12;
|
||
D = 5*L; % traceback (unused here; full TB)
|
||
Nwin = 7; % BM feature window
|
||
Delta = 0;
|
||
|
||
%% TX, channel, noise
|
||
map = @(b) 2*b-1;
|
||
b_tr = randi([0 1],Ktr,1); x_tr = map(b_tr);
|
||
b_te = randi([0 1],Kte,1); x_te = map(b_te);
|
||
|
||
convpad = @(x) filter(h,1,[x; zeros(M-1,1)]);
|
||
ytr_clean = convpad(x_tr);
|
||
yte_clean = convpad(x_te);
|
||
|
||
EsN0 = 10^(EsN0dB/10);
|
||
sigma2 = 1/(2*EsN0);
|
||
awgn = @(len) sqrt(sigma2)*randn(len,1);
|
||
y_tr = ytr_clean + awgn(length(ytr_clean));
|
||
y_te = yte_clean + awgn(length(yte_clean));
|
||
|
||
% remove tail to keep equal lengths
|
||
trim = M-1;
|
||
y_tr = y_tr(1:end-trim); x_tr = x_tr(1:end-trim); b_tr = b_tr(1:end-trim);
|
||
y_te = y_te(1:end-trim); x_te = x_te(1:end-trim); b_te = b_te(1:end-trim);
|
||
|
||
%% Trellis
|
||
nStates = 2^(L-1);
|
||
statesBin = de2bi(0:nStates-1,L-1,'left-msb');
|
||
symOfBit = @(bit) (2*bit-1);
|
||
|
||
% transitions (s'->s) with input symbol u in {±1}
|
||
trans = struct('sp',[],'s',[],'u',[],'ubit',[]);
|
||
idx=1;
|
||
for sp=1:nStates
|
||
prev_bits = statesBin(sp,:); % [x_{k-1}, x_{k-2}] as bits
|
||
for ubit=[0 1]
|
||
u = symOfBit(ubit);
|
||
nb = [prev_bits(1),ubit]; % new state s = [x_k, x_{k-1}]
|
||
s = bi2de(nb,'left-msb')+1;
|
||
trans(idx).sp=sp; trans(idx).s=s; trans(idx).u=u; trans(idx).ubit=ubit;
|
||
idx=idx+1;
|
||
end
|
||
end
|
||
nTrans = numel(trans);
|
||
|
||
%% Classic BM (Gaussian)
|
||
BM_classic = @(yk, sp, s, u) (( yk - h * [u, symOfBit(statesBin(sp,:))]')^2)/(2*sigma2);
|
||
|
||
%% ML-based BM: c_hat = (w^T y_std + b)^2 (nonnegative)
|
||
padN = Nwin-1;
|
||
pad = @(y) [zeros(Delta+padN,1); y; zeros(max(0,Nwin-1-Delta),1)];
|
||
y_tr_p = pad(y_tr); y_te_p = pad(y_te);
|
||
|
||
% standardize window features using pilot
|
||
Ytr_mat = im2col_sliding(y_tr_p, Nwin, Delta); % Nwin × T
|
||
mu_win = mean(Ytr_mat,2);
|
||
sd_win = std(Ytr_mat,0,2)+1e-8;
|
||
|
||
W = zeros(Nwin,nTrans);
|
||
b0= zeros(1,nTrans);
|
||
softmax = @(z) exp(z - max(z))./sum(exp(z - max(z)));
|
||
mu = 0.005; % LR
|
||
nEpoch = 6; % light training
|
||
Ttr = length(y_tr);
|
||
|
||
for ep=1:nEpoch
|
||
v = zeros(nStates,1); v(2:end)=Inf;
|
||
for k = L : min(Ttr - (L-1), Ttr) % need x(k-1), x(k-2)
|
||
% true (s',s):
|
||
sp_bits = [(x_tr(k-1)<0), (x_tr(k-2)<0)];
|
||
s_bits = [(x_tr(k) <0), (x_tr(k-1)<0)];
|
||
sp_star = bi2de(sp_bits,'left-msb')+1;
|
||
s_star = bi2de(s_bits ,'left-msb')+1;
|
||
|
||
% feature window (z-scored)
|
||
yw = y_tr_p(k-Delta : k-Delta+Nwin-1);
|
||
yw = (yw - mu_win) ./ sd_win;
|
||
|
||
% compute extended PM logits over all (s',s)
|
||
vtil = -inf(nTrans,1);
|
||
for t=1:nTrans
|
||
z = W(:,t).'*yw + b0(t);
|
||
c_hat = z*z; % (.)^2 nonnegative BM
|
||
vtil(t) = -( v(trans(t).sp) + c_hat);
|
||
end
|
||
pext = softmax(vtil).';
|
||
t_star = find([trans.sp]==sp_star & [trans.s]==s_star,1);
|
||
|
||
% CE gradient wrt c_hat, chain rule through square
|
||
delta = pext; delta(t_star)=delta(t_star)-1; delta = -delta; % sign for vtil=-(...)
|
||
for t=1:nTrans
|
||
z = W(:,t).'*yw + b0(t);
|
||
dc_dz = 2*z; % d(z^2)/dz
|
||
g = delta(t) * dc_dz; % ∂L/∂z
|
||
W(:,t) = W(:,t) - mu * g * yw;
|
||
b0(t) = b0(t) - mu * g;
|
||
end
|
||
|
||
% PM update for next step
|
||
v_next = inf(nStates,1);
|
||
for t=1:nTrans
|
||
sp=trans(t).sp; s=trans(t).s;
|
||
z = W(:,t).'*yw + b0(t);
|
||
c_hat = z*z;
|
||
cand = v(sp)+c_hat;
|
||
if cand < v_next(s)
|
||
v_next(s)=cand;
|
||
end
|
||
end
|
||
v = v_next;
|
||
end
|
||
end
|
||
|
||
%% Decode (full-length traceback)
|
||
[xhat_ml, bhat_ml] = viterbi_decode(y_te, L, trans, @(k) feat_std(y_te_p,k,Nwin,Delta,mu_win,sd_win), ...
|
||
@(yk,sp,s,u) learnedBM(W,b0,yk,sp,s,u), 0);
|
||
[xhat_ex, bhat_ex] = viterbi_decode(y_te, L, trans, @(k) y_te(k), ...
|
||
@(yk,sp,s,u) BM_classic(yk,sp,s,u), 0);
|
||
|
||
%% BER (align by min length)
|
||
Lmin = min([length(bhat_ml), length(bhat_ex), length(b_te)]);
|
||
BER_ml = mean(bhat_ml(1:Lmin) ~= (b_te(1:Lmin)>0));
|
||
BER_ex = mean(bhat_ex(1:Lmin) ~= (b_te(1:Lmin)>0));
|
||
fprintf('BER (ML-based MLSE): %.3e\n', BER_ml);
|
||
fprintf('BER (classic MLSE ): %.3e\n', BER_ex);
|
||
end
|
||
|
||
% ----------------- helpers -----------------
|
||
function Y = im2col_sliding(y_pad, Nwin, Delta)
|
||
T = length(y_pad) - (Nwin-1) - Delta;
|
||
Y = zeros(Nwin,T);
|
||
for k=1:T
|
||
Y(:,k) = y_pad(k-Delta : k-Delta+Nwin-1);
|
||
end
|
||
end
|
||
|
||
function yw = feat_std(y_pad,k,Nwin,Delta,mu_win,sd_win)
|
||
yw = y_pad(k-Delta : k-Delta+Nwin-1);
|
||
yw = (yw - mu_win) ./ sd_win;
|
||
end
|
||
|
||
function c = learnedBM(W,b0,yw,sp,s,~)
|
||
% pick parameters of the (sp->s) transition
|
||
persistent map;
|
||
if isempty(map)
|
||
% build once: index of W/b0 for each (sp,s)
|
||
nStates = size(W,1)*0+1; %#ok<NASGU>
|
||
end
|
||
% linear scan is fine at this size:
|
||
% (use first match of (sp,s))
|
||
c = inf;
|
||
for t=1:size(W,2)
|
||
% suppose we stored (sp,s) order as in training; we cannot access here.
|
||
% Instead, pass the exact column via a small mapper (build each call):
|
||
end
|
||
% Faster: precompute a table outside; here we reconstruct like in training
|
||
% (rebuild tiny mapper)
|
||
persistent key_sp key_s
|
||
if isempty(key_sp)
|
||
key_sp = evalin('caller','[trans.sp]');
|
||
key_s = evalin('caller','[trans.s]');
|
||
end
|
||
t = find(key_sp==sp & key_s==s,1);
|
||
z = W(:,t).'*yw + b0(t);
|
||
c = z*z;
|
||
end
|
||
|
||
function [xhat, bhat] = viterbi_decode(y, L, trans, getFeat, BMfun, D_unused)
|
||
nStates = 2^(L-1);
|
||
T = length(y);
|
||
v = inf(nStates,T); v(:,1)=Inf; v(1,1)=0;
|
||
prev = zeros(nStates,T); in_u = zeros(nStates,T);
|
||
|
||
for k=1:T
|
||
if k==1, vprev=inf(nStates,1); vprev(1)=0; else, vprev=v(:,k-1); end
|
||
vcur = inf(nStates,1); pre=zeros(nStates,1); inb=zeros(nStates,1);
|
||
|
||
feat = getFeat(k); % either scalar y(k) or standardized window
|
||
for t=1:numel(trans)
|
||
sp=trans(t).sp; s=trans(t).s; u=trans(t).u;
|
||
ck = BMfun(feat, sp, s, u);
|
||
cand = vprev(sp)+ck;
|
||
if cand < vcur(s)
|
||
vcur(s)=cand; pre(s)=sp; inb(s)=u;
|
||
end
|
||
end
|
||
v(:,k)=vcur; prev(:,k)=pre; in_u(:,k)=inb;
|
||
end
|
||
|
||
[~,st]=min(v(:,T));
|
||
xhat=zeros(T,1);
|
||
for k=T:-1:1
|
||
xhat(k)=in_u(st,k);
|
||
st=prev(st,k); if st==0, st=1; end
|
||
end
|
||
bhat = xhat>0;
|
||
end
|