Current work is on MLSE and SD Decoding etc. MLSE is currently not 100% working, the scalings are maybe off?!
153 lines
4.9 KiB
Matlab
153 lines
4.9 KiB
Matlab
classdef MLSE_new < handle
|
||
%MLSE: BCJR‐based soft‐output for PAM‐M sequence estimation & GMI
|
||
|
||
properties
|
||
M % PAM order (e.g. 4)
|
||
DIR % channel impulse response
|
||
trellis_states % PAM constellation levels (e.g. [-3 -1 1 3])
|
||
end
|
||
|
||
methods
|
||
function obj = MLSE_new(opts)
|
||
arguments
|
||
opts.M double = 4;
|
||
opts.DIR double = 1;
|
||
opts.trellis_states double = [-3 -1 1 3];
|
||
end
|
||
obj.M = opts.M;
|
||
obj.DIR = opts.DIR;
|
||
obj.trellis_states = opts.trellis_states;
|
||
end
|
||
|
||
function [hd_out, LLR, NGMI] = process(obj, signalclass,ref_symbolclass)
|
||
% rx, tx: column vectors of equalized and reference symbols
|
||
data_in = signalclass.signal;
|
||
shape_in = size(data_in);
|
||
data_ref = ref_symbolclass.signal;
|
||
|
||
[data_out_hd,LLR,GMI] = obj.bcjr_soft(data_in,data_ref);
|
||
try
|
||
data_out_hd = reshape(data_out_hd,shape_in(1),shape_in(2));
|
||
catch
|
||
warning('output reshaping failed after MLSE');
|
||
end
|
||
|
||
signalclass_hd = signalclass;
|
||
signalclass_hd.signal = data_out_hd;
|
||
|
||
|
||
|
||
|
||
[hd_out, LLR, NGMI] = obj.bcjr_soft(rx, tx);
|
||
end
|
||
end
|
||
|
||
methods (Access=private)
|
||
function [hd_sym, LLR, NGMI] = bcjr_soft(obj, y, x)
|
||
%--- build trellis states & branch outputs ---
|
||
DIR = flip(obj.DIR(:));
|
||
S = combvec(obj.trellis_states, obj.trellis_states)';
|
||
prev = S(:,1); cur = S(:,2);
|
||
branch_out = prev*DIR(1) + cur*DIR(2); % for 2‐tap
|
||
|
||
N = numel(y);
|
||
M = obj.M;
|
||
m = log2(M);
|
||
|
||
%--- forward/backward metrics (max‐log) ---
|
||
nStates = numel(obj.trellis_states);
|
||
alpha = -inf(nStates,N);
|
||
beta = -inf(nStates,N);
|
||
bm_fw = zeros(nStates,nStates,N);
|
||
|
||
% branch metrics
|
||
sigma2 = var(y - x);
|
||
inv2sigma2 = 1/(2*sigma2);
|
||
for n=1:N
|
||
bm = -((y(n)-branch_out).^2)* inv2sigma2;
|
||
bm = reshape(bm, nStates, nStates);
|
||
bm_fw(:,:,n) = bm;
|
||
end
|
||
% forward
|
||
alpha(:,1) = max(bm_fw(:,:,1),[],2);
|
||
for n=2:N
|
||
mat = alpha(:,n-1) + bm_fw(:,:,n);
|
||
alpha(:,n) = max(mat,[],2);
|
||
end
|
||
% backward
|
||
beta(:,N) = 0;
|
||
for n=N-1:-1:1
|
||
mat = beta(:,n+1).' + bm_fw(:,:,n+1);
|
||
beta(:,n) = max(mat,[],2);
|
||
end
|
||
|
||
% Precompute the “zero‐th” forward metric
|
||
alpha0 = zeros(nStates,1);
|
||
|
||
LLP = -inf(nStates,nStates,N);
|
||
for n = 1:N
|
||
% Select the correct previous alpha column
|
||
if n == 1
|
||
a_prev = alpha0;
|
||
else
|
||
a_prev = alpha(:,n-1);
|
||
end
|
||
|
||
% Combine α, branch metric, and β
|
||
mat = a_prev + bm_fw(:,:,n) + beta(:,n).';
|
||
LLP(:,:,n) = mat;
|
||
end
|
||
|
||
%--- compute LLRs symbol‐wise via log‐sum‐exp over branches ---
|
||
levels = obj.trellis_states;
|
||
bit_map = PAMmapper(M,0).demap(levels');
|
||
LLR = zeros(N,m);
|
||
for n=1:N
|
||
% sum branch posteriors per symbol
|
||
llp_n = LLP(:,:,n);
|
||
logZ = logsumexp(llp_n(:));
|
||
post = exp(llp_n - logZ);
|
||
% aggregate per symbol index
|
||
Psym = zeros(M,1);
|
||
for i=1:nStates
|
||
for j=1:nStates
|
||
% branch (i->j) emits symbol 'cur(j)'
|
||
idx = find(levels==cur(j));
|
||
Psym(idx) = Psym(idx) + post(i,j);
|
||
end
|
||
end
|
||
% bit‐LLR from symbol posterior
|
||
for k=1:m
|
||
idx1 = bit_map(:,k)==1;
|
||
idx0 = ~idx1;
|
||
P1 = sum(Psym(idx1));
|
||
P0 = sum(Psym(idx0));
|
||
LLR(n,k) = log(P1/P0);
|
||
end
|
||
end
|
||
|
||
%--- hard decisions ---
|
||
[~,symIdx] = max(LLR,[],2);
|
||
hd_sym = levels(symIdx);
|
||
|
||
%--- GMI via Alvarado Eq.(30) ---
|
||
tx_bits = PAMmapper(M,0).demap(x);
|
||
MI = zeros(1,m);
|
||
for k=1:m
|
||
r0 = LLR(tx_bits(:,k)==0,k);
|
||
r1 = LLR(tx_bits(:,k)==1,k);
|
||
I0 = mean(log2(1+exp(-r0)));
|
||
I1 = mean(log2(1+exp(+r1)));
|
||
MI(k) = 1 - 0.5*(I0+I1);
|
||
end
|
||
GMI = sum(MI);
|
||
NGMI = GMI/m;
|
||
end
|
||
end
|
||
end
|
||
|
||
function s = logsumexp(a)
|
||
m = max(a(:));
|
||
s = m + log(sum(exp(a-m), 'all'));
|
||
end
|