Files
imdd_silas/projects/ML_based_MLSE/minimal_model_gpt.m
Silas Oettinghaus 09f345aa91 ML_MLSE (before testing in depth)
minor changes here and there
2025-10-31 10:36:13 +01:00

214 lines
6.7 KiB
Matlab
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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