Files
imdd_silas/Classes/04_DSP/Equalizer/ML_MLSE.m
2025-10-27 07:52:51 +01:00

274 lines
10 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.
classdef ML_MLSE < handle
% Implementation of plain and simple FFE.
% 1) Training mode (stable performance when you use NLMS)
% 2) Decision directed mode
%LMS: mu in order of 0.0001 for acceptable convergence speed
%NLMS: mu in order of 0.01 for acceptable convergence speed
%RLS: mu is lambda -> 0.99 -> 1 (has a strong dependency on this! use a loop to find out best values)
% FFE("epochs_tr",5,"epochs_dd",2,"len_tr",2^13,"mu_dd",mu_dd,"mu_tr",mu_tr,"order",25,"sps",2,"decide",0, "adaption",adaption_method(adaption),"dd_mode",use_dd_mode);
properties
sps % usually 2
order
e
e_tr
error
len_tr
mu_tr
epochs_tr
dd_mode % 1 or 0 to set DD-mode on or off
mu_dd %weight update in dd mode
epochs_dd
constellation
L %viterbi memory length
alpha
DIR
DIR_flip
trellis_states
traceback_depth
end
methods
function obj = ML_MLSE(options)
arguments(Input)
options.sps = 2;
options.order = 15;
options.len_tr = 4096;
options.mu_tr = 0;
options.epochs_tr = 5;
options.dd_mode = 1;
options.mu_dd = 1e-5;
options.epochs_dd = 5;
options.traceback_depth = 1024;
options.L = 1
end
fn = fieldnames(options);
for n = 1:numel(fn)
obj.(fn{n}) = options.(fn{n});
end
obj.e = zeros(obj.order,1);
obj.error = 0;
end
function [X,X_viterbi] = process(obj, X, D)
% actual processing of the signal (steps 1. - 3.)
% 1 normalize RMS
X = X.normalize("mode","rms");
obj.constellation = unique(D.signal);
if length(X)/length(D) ~= obj.sps
warning('Signal length does not fit to reference!');
end
% Training Mode
n = obj.len_tr;
training = 1;
obj.equalize(X.signal, D.signal,obj.mu_tr,obj.epochs_dd,n,training);
obj.e_tr = obj.e;
% Decision Directed Mode
n = X.length;
training = 0;
[y,y_vit]=obj.equalize(X.signal, D.signal,obj.mu_dd,obj.epochs_dd,n,training);
X_viterbi = X;
X.signal = y;
X.fs = D.fs; %change sampling frequency of outgoing signal from fdac e.g. 2 sps to symbol spaced = fsym
lbdesc = [num2str(obj.order),' tap FFE'];
X = X.logbookentry(lbdesc); % append to logbook
X_viterbi.signal = y_vit;
X_viterbi.fs = D.fs; %change sampling frequency of outgoing signal from fdac e.g. 2 sps to symbol spaced = fsym
lbdesc = [num2str(obj.order),'order FFE + PF + Viterbi'];
X_viterbi = X_viterbi.logbookentry(lbdesc); % append to logbook
end
function [y,y_vit] = equalize(obj,x,d,mu,epochs,N,training)
% ==============================================================
% FFE + Whitening + ML-Based Branch Metric Estimation + Viterbi
% ==============================================================
% --- Input padding and preallocation
x = [zeros(floor(obj.order/2),1); x; zeros(obj.order,1)];
N_ = N / obj.sps;
y = zeros(N_,1);
y_white = zeros(N_,1);
for epoch = 1:epochs
% ==============================================================
% INITIALIZATION (only before final epoch and detection mode)
% ==============================================================
if epoch == epochs
% --- Parameters
S = numel(unique(d)); % alphabet size
Nf = 9; % filter length
Delta = ceil(Nf/2); % delay parameter
nStates = S^obj.L;
nFeasible = nStates*S;
% --- Trellis mapping
obj.DIR = arburg(y-d(1:N_), obj.L);
obj.DIR_flip = flip(obj.DIR);
obj.trellis_states = reshape(unique(d),1,[]);
pre_comb_mat = repmat(obj.trellis_states, obj.L, 1);
pre_comb_cell = mat2cell(pre_comb_mat, ones(1,obj.L), size(pre_comb_mat,2));
combs = fliplr(combvec(pre_comb_cell{:}).');
first_sym = combs(:,1);
last_sym = combs(:,end);
nStates = size(combs,1);
% --- Valid transitions
valid = false(nStates);
for from = 1:nStates
for to = 1:nStates
if all(combs(to,2:end) == combs(from,1:end-1))
valid(to,from) = true;
end
end
end
[valid_to, valid_from] = find(valid);
% --- Noise estimation
y_ideal = conv(d(1:N_), obj.DIR(:), "same");
sigma2 = mean(abs(y - y_ideal).^2);
inv2s2 = 1/(2*sigma2);
% --- Allocate vectors and weights
pm = zeros(nStates,1);
w = zeros(Nf,nFeasible); % filter weights per transition
b = zeros(1,nFeasible); % bias terms
v_hat = zeros(1,nFeasible);
v_tilde = zeros(1,nFeasible);
bm_vec = zeros(1,nFeasible);
zi = zeros(max(numel(obj.DIR)-1,0),1);
mu_w = 0.001;
mu_b = 0.001;
end
% ==============================================================
% RUNTIME LOOP
% ==============================================================
symbol = 0;
for sample = 1:obj.sps:N
symbol = symbol + 1;
% --- FFE output
U = x(obj.order+sample-1:-1:sample);
y(symbol,1) = obj.e.' * U;
% --- Decision / FFE adaptation
if training
d_hat = d(symbol);
else
[~,idx] = min(abs(y(symbol) - obj.constellation));
d_hat = obj.constellation(idx);
end
err = d_hat - y(symbol);
obj.e = obj.e + mu * (err * U);
% --- Whitening + MLSE in last epoch
if epoch == epochs
[y_white(symbol), zi] = filter(obj.DIR,1,y(symbol), zi);
k = symbol;
% --- Build Δ-delayed observation window y_k
i1 = k - Nf + 1 + Delta;
i2 = k + Delta;
buf = y_white(max(1,i1):min(length(y_white),i2));
padL = max(0,1 - i1);
padR = max(0,i2 - length(y_white));
yk = [zeros(padL,1); buf(:); zeros(padR,1)]; % Nf×1
% --- Predict branch metrics for all feasible transitions
v_hat = (yk.' * w) + b; % [1×nFeasible]
v_hat = v_hat.'; % [nFeasible×1]
% --- Extended path metrics
v_tilde = pm(valid_from) + v_hat; % [nFeasible×1]
% ===== Gradient update (Algorithm 1) =====
% if training
% for current symbol index k -> previous (k-1) and current (k)
if k > obj.L
prev_seq = d(k-obj.L:k-1); % previous state symbols
curr_seq = d(k-obj.L+1:k); % next state symbols
% find state indices in trellis
true_from_state = find(ismember(combs, prev_seq.', 'rows'));
true_to_state = find(ismember(combs, curr_seq.', 'rows'));
else
% not enough history yet
true_from_state = 1;
true_to_state = 1;
end
% softmax over -v_tilde, one-hot target t
p = exp(-v_tilde);
p = p./sum(p);
t = zeros(nFeasible,1);
t(valid_from==true_from_state & valid_to==true_to_state) = 1;
delta = t - p; % ∂CE/∂(-v_tilde)
w = w + mu_w * (yk * delta.'); % Nf×nFeasible
b = b + mu_b * delta.'; % 1×nFeasible
% end
% =========================================
% compareselect to next states
pm_next = -inf(nStates,1);
surv_idx = zeros(nStates,1);
for s_to = 1:nStates
mask = (valid_to==s_to);
[pm_next(s_to), arg] = max(-v_tilde(mask)); % max likelihood ↔ min metric
surv_idx(s_to) = valid_from(find(mask,1,'first')-1+arg);
end
pm = pm_next;
% --- Traceback
if mod(symbol,obj.traceback_depth) == 0
[~,viterbi_path(symbol)] = max(pm);
for n = symbol:-1:symbol-obj.traceback_depth+2
viterbi_path(n-1) = surv_idx(viterbi_path(n));
end
end
end
end
% --- Output reconstructed path
if epoch == epochs && ~training
y_vit = first_sym(viterbi_path);
end
end
end
end
end