Developing the ML-enhanced MLSE :-)

This commit is contained in:
silas (home)
2025-10-27 07:52:51 +01:00
parent 7085ba0931
commit ef4dc53db5
3 changed files with 57 additions and 43 deletions

View File

@@ -1,4 +1,4 @@
classdef ML_MLSE < handle classdef FFE_MLSE < handle
% Implementation of plain and simple FFE. % Implementation of plain and simple FFE.
% 1) Training mode (stable performance when you use NLMS) % 1) Training mode (stable performance when you use NLMS)
% 2) Decision directed mode % 2) Decision directed mode
@@ -37,7 +37,7 @@ classdef ML_MLSE < handle
end end
methods methods
function obj = ML_MLSE(options) function obj = FFE_MLSE(options)
arguments(Input) arguments(Input)
options.sps = 2; options.sps = 2;

View File

@@ -121,22 +121,22 @@ classdef ML_MLSE < handle
% ============================================================== % ==============================================================
% INITIALIZATION (only before final epoch and detection mode) % INITIALIZATION (only before final epoch and detection mode)
% ============================================================== % ==============================================================
if epoch == epochs && ~training if epoch == epochs
% --- Parameters % --- Parameters
S = numel(unique(d)); % alphabet size S = numel(unique(d)); % alphabet size
L = obj.L; % MLSE memory Nf = 9; % filter length
Nf = L; % filter length Delta = ceil(Nf/2); % delay parameter
Delta = ceil(L/2); % delay parameter nStates = S^obj.L;
nStates = S^L; nFeasible = nStates*S;
nFeasible = S^(L-1)*S;
% --- Trellis mapping % --- Trellis mapping
obj.DIR = arburg(y-d, L); obj.DIR = arburg(y-d(1:N_), obj.L);
obj.DIR_flip = flip(obj.DIR); obj.DIR_flip = flip(obj.DIR);
obj.trellis_states = reshape(unique(d),1,[]); obj.trellis_states = reshape(unique(d),1,[]);
pre_comb_mat = repmat(obj.trellis_states, L, 1); pre_comb_mat = repmat(obj.trellis_states, obj.L, 1);
pre_comb_cell = mat2cell(pre_comb_mat, ones(1,L), size(pre_comb_mat,2)); pre_comb_cell = mat2cell(pre_comb_mat, ones(1,obj.L), size(pre_comb_mat,2));
combs = fliplr(combvec(pre_comb_cell{:}).'); combs = fliplr(combvec(pre_comb_cell{:}).');
first_sym = combs(:,1); first_sym = combs(:,1);
last_sym = combs(:,end); last_sym = combs(:,end);
@@ -154,7 +154,7 @@ classdef ML_MLSE < handle
[valid_to, valid_from] = find(valid); [valid_to, valid_from] = find(valid);
% --- Noise estimation % --- Noise estimation
y_ideal = conv(d(:), obj.DIR(:), "same"); y_ideal = conv(d(1:N_), obj.DIR(:), "same");
sigma2 = mean(abs(y - y_ideal).^2); sigma2 = mean(abs(y - y_ideal).^2);
inv2s2 = 1/(2*sigma2); inv2s2 = 1/(2*sigma2);
@@ -166,6 +166,9 @@ classdef ML_MLSE < handle
v_tilde = zeros(1,nFeasible); v_tilde = zeros(1,nFeasible);
bm_vec = zeros(1,nFeasible); bm_vec = zeros(1,nFeasible);
zi = zeros(max(numel(obj.DIR)-1,0),1); zi = zeros(max(numel(obj.DIR)-1,0),1);
mu_w = 0.001;
mu_b = 0.001;
end end
% ============================================================== % ==============================================================
@@ -190,7 +193,7 @@ classdef ML_MLSE < handle
obj.e = obj.e + mu * (err * U); obj.e = obj.e + mu * (err * U);
% --- Whitening + MLSE in last epoch % --- Whitening + MLSE in last epoch
if epoch == epochs && ~training if epoch == epochs
[y_white(symbol), zi] = filter(obj.DIR,1,y(symbol), zi); [y_white(symbol), zi] = filter(obj.DIR,1,y(symbol), zi);
k = symbol; k = symbol;
@@ -210,17 +213,42 @@ classdef ML_MLSE < handle
% --- Extended path metrics % --- Extended path metrics
v_tilde = pm(valid_from) + v_hat; % [nFeasible×1] v_tilde = pm(valid_from) + v_hat; % [nFeasible×1]
% --- Compute branch metrics (distance) % ===== Gradient update (Algorithm 1) =====
bm_vec = -(y_white(k) - v_hat).^2 * inv2s2; % 1×nFeasible % if training
% --- Survivor selection (vector aggregation) % for current symbol index k -> previous (k-1) and current (k)
pm_new_vec = pm(valid_from) + bm_vec.'; % nFeasible×1 if k > obj.L
pm_next = -inf(nStates,1); prev_seq = d(k-obj.L:k-1); % previous state symbols
surv_idx = zeros(nStates,1); curr_seq = d(k-obj.L+1:k); % next state symbols
for t = 1:nStates
mask = (valid_to==t); % find state indices in trellis
[pm_next(t), arg] = max(pm_new_vec(mask)); true_from_state = find(ismember(combs, prev_seq.', 'rows'));
surv_idx(t) = valid_from(find(mask,1,'first')-1+arg); 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 end
pm = pm_next; pm = pm_next;

View File

@@ -99,15 +99,15 @@ for r = 1:length(fsym)
%% Implement DSP directly here: %% Implement DSP directly here:
mu_lms = 0.0005; mu_lms = 0.0005;
% tic
% eq = FFE_MLSE("epochs_tr",5,"epochs_dd",5,"len_tr",2^13,"mu_dd",mu_lms,"mu_tr",mu_lms,"order",50,"sps",2,"traceback_depth",32,"L",1);
% [Eq_signal,Vit_signal] = eq.process(Rx_sig_2sps,Symbols);
% toc
tic tic
eq = ML_MLSE("epochs_tr",5,"epochs_dd",5,"len_tr",2^13,"mu_dd",mu_lms,"mu_tr",mu_lms,"order",50,"sps",2,"traceback_depth",32,"L",1); eq = ML_MLSE("epochs_tr",5,"epochs_dd",5,"len_tr",2^13,"mu_dd",mu_lms,"mu_tr",mu_lms,"order",50,"sps",2,"traceback_depth",32,"L",1);
[Eq_signal,Vit_signal] = eq.process(Rx_sig_2sps,Symbols); [Eq_signal,Vit_signal] = eq.process(Rx_sig_2sps,Symbols);
toc toc
% tic
% eq = FFE_MLSE("epochs_tr",5,"epochs_dd",5,"len_tr",2^13,"mu_dd",mu_lms,"mu_tr",mu_lms,"order",50,"sps",2,"traceback_depth",32,"L",3);
% [Eq_signal,Vit_signal] = eq.process(Rx_sig_2sps,Symbols);
% toc
Eq_bits = PAMmapper(M, 0, "eth_style", 0).demap(Eq_signal); Eq_bits = PAMmapper(M, 0, "eth_style", 0).demap(Eq_signal);
[~, errors, ber, ~] = calc_ber(Eq_bits.signal, Tx_bits.signal, "skip_front", 0, "skip_end", 0, "returnErrorLocation", 1); [~, errors, ber, ~] = calc_ber(Eq_bits.signal, Tx_bits.signal, "skip_front", 0, "skip_end", 0, "returnErrorLocation", 1);
@@ -122,12 +122,11 @@ for r = 1:length(fsym)
%% optimize smth. %% optimize smth.
tr_len = 2.^[2:15]; tr_len = 2.^[2:15];
tr_len = floor(tr_len); tr_len = floor(tr_len);
ber = zeros(size(tr_len)); ber = zeros(size(tr_len));
parfor m = 1:numel(tr_len) for m = 1:numel(tr_len)
mu_lms = 0.0005; mu_lms = 0.0005;
eq = ML_MLSE("epochs_tr",5,"epochs_dd",5,"len_tr",2^13,"mu_dd",mu_lms,"mu_tr",mu_lms,"order",50,"sps",2,"traceback_depth",tr_len(m)); eq = ML_MLSE("epochs_tr",5,"epochs_dd",5,"len_tr",2^13,"mu_dd",mu_lms,"mu_tr",mu_lms,"order",50,"sps",2,"traceback_depth",tr_len(m));
@@ -150,19 +149,6 @@ for r = 1:length(fsym)
set(gca,'YScale','log'); set(gca,'YScale','log');
%% RUN Comparison %% RUN Comparison
len_tr = 4096*2; len_tr = 4096*2;