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

364 lines
14 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
% --- Added internal class variables used later ---
S
Nf
delta
nStates
nFeasible
combs
first_sym
last_sym
valid
valid_to_idx
valid_from_idx
w
% --- New: fast state lookup ---
state_dict % containers.Map: key(sequence)->state index
key_fmt = '%.8g_'; % key format for sequence strings
nSym % |constellation|
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.delta = 0;
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");
% Use sorted constellation for deterministic mapping
obj.constellation = sort(unique(D.signal),'ascend');
obj.nSym = numel(obj.constellation);
if length(X)/length(D) ~= obj.sps
warning('Signal length does not fit to reference!');
end
% ==============================================================
% INITIALIZATION (only before final epoch and detection mode)
% ==============================================================
% --- Parameters
obj.S = numel(obj.constellation); % alphabet size
obj.Nf = obj.order*obj.sps; % filter length
% obj.delta = 3;%ceil(obj.Nf/2); % delay parameter
obj.nStates = obj.S^obj.L;
obj.nFeasible = obj.nStates*obj.S;
% --- Trellis mapping
obj.trellis_states = reshape(obj.constellation,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));
obj.combs = fliplr(combvec(pre_comb_cell{:}).'); % rows: states, columns: [x_k, x_{k-1}, ...]
obj.first_sym = obj.combs(:,1);
obj.last_sym = obj.combs(:,end);
obj.nStates = size(obj.combs,1);
% --- Valid transitions
obj.valid = false(obj.nStates);
for from = 1:obj.nStates
for to = 1:obj.nStates
if all(obj.combs(to,2:end) == obj.combs(from,1:end-1))
obj.valid(to,from) = true;
end
end
end
[obj.valid_to_idx, obj.valid_from_idx] = find(obj.valid);
% --- Allocate vectors and weights
% !! IF SHAPE FIT, then we already have smth there an we want
% to start with the existing fitler-set
if isempty(obj.w) || any(size(obj.w) ~= [obj.Nf+1,obj.nFeasible])
obj.w = zeros(obj.Nf+1,obj.nFeasible); % filter weights per transition + bias tap
obj.w = randn(obj.Nf+1,obj.nFeasible);
end
% --- Precompute dictionary for fast state lookup (sequence -> state)
keys = cell(obj.nStates,1);
for i = 1:obj.nStates
keys{i} = obj.seq_key(obj.combs(i,:)); % combs row is already [x_k, x_{k-1}, ...]
end
obj.state_dict = containers.Map(keys, 1:obj.nStates);
% ==============================================================
% TRAINING
% ==============================================================
% Training Mode
n = obj.len_tr;
training = 1;
obj.equalize(X.signal, D.signal,obj.mu_tr,obj.epochs_tr,n,training);
obj.e_tr = obj.e;
% ==============================================================
% DD-Mode / Fixed Mode
% ==============================================================
% 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
% ==============================================================
debug = 0;
% --- Input padding and preallocation
y = zeros(N,1);
% number of symbol steps in this block
nSymbols = ceil(N/obj.sps);
for epoch = 1:epochs
% state metrics (log-domain costs): keep as column [nStates×1]
pm = zeros(obj.nStates,1); % v_{k-1}(s)
c_hat = zeros(1,obj.nFeasible);
v_tilde = zeros(1,obj.nFeasible);
pred = zeros(nSymbols, obj.nStates, 'uint32');
pm_sto = nan(obj.nStates, nSymbols,'like',pm);
% --- Initialize "true" trellis state for training (shift-register style)
% expect sequences in chronological order [x_{k-L+1}:x_k], but combs rows are [x_k, x_{k-1}, ...]
if numel(d) >= obj.L
init_seq = d(1:obj.L); % [x_1 ... x_L]
true_to_state_idx = obj.state_dict(obj.seq_key(flip(init_seq))); % flip to [x_L, x_{L-1}, ...]
else
true_to_state_idx = 1;
end
symbol = 0;
for sample = 1:obj.sps:N
symbol = symbol + 1;
k = symbol;
% --- Build Δ-delayed observation window y_k
i1 = sample - obj.Nf + 1 + obj.delta;
i2 = sample + obj.delta;
buf = x(max(1,i1):min(length(x),i2));
padL = max(0,1 - i1);
padR = max(0,i2 - length(x));
yk = [zeros(padL,1); buf(:); zeros(padR,1)]; % Nf×1
yk = [yk;1];
% --- Predict branch metrics for all feasible transitions: c_hat
c_hat = (yk.' * obj.w); % [1×nFeasible]
c_hat = c_hat.'; % [nFeasible×1]
% --- Extended path metrics: v_tilde = pm(from) + c_hat
% normalize pm to avoid growth (invariant to additive const)
pm = pm - min(pm);
v_tilde = pm(obj.valid_from_idx) + c_hat; % [nFeasible×1]
% ===== Gradient update (Algorithm 1) =====
% if training
if k > obj.L
% shift-register: previous "to" becomes current "from"
true_from_state_idx = true_to_state_idx;
% current "to" from data window
curr_seq = d(k-obj.L+1:k);
key_to = obj.seq_key(flip(curr_seq));
if isKey(obj.state_dict, key_to)
true_to_state_idx = obj.state_dict(key_to);
else
% fall back safely (should not happen with proper constellation)
true_to_state_idx = true_from_state_idx;
end
else
% not enough history yet
true_from_state_idx = 1;
true_to_state_idx = true_to_state_idx; % keep init
end
if 0
disp(['FROM: state',char(num2str(true_from_state_idx)),' : symbol transition', char(num2str(obj.combs(true_from_state_idx,:)))]);
disp(['TO: state',char(num2str(true_to_state_idx)),' : symbol transition', char(num2str(obj.combs(true_to_state_idx,:)))]);
end
% Dirac delta over correct extended transition (from,to)
dirac = zeros(obj.nFeasible,1);
dirac(obj.valid_from_idx==true_from_state_idx & obj.valid_to_idx==true_to_state_idx) = 1; % This Dirac delta function δ(s = s k, s = s k1) = 1 if the extended state (s, s) corresponds to the true realized states (s k, s k1), and is zero otherwise.
% softmax over -v_tilde (numerically safe shift)
p = exp(-(v_tilde - max(v_tilde)));
p = p./sum(p); % found in formula (9) and (19)
% gradient term (t - p)
dmp = (dirac - p)'; % 1×nFeasible
if mod(symbol,128) == 1 && debug
% --- Normalize and compute probabilities
v_norm = v_tilde - max(v_tilde);
probs_lin = exp(-v_norm); probs_lin = probs_lin ./ sum(probs_lin);
probs_log = -v_norm;
% --- Map back into nStates×nStates grid
probs_mat = nan(obj.nStates, obj.nStates);
probs_logmat = nan(obj.nStates, obj.nStates);
probs_mat(obj.valid) = probs_lin;
probs_logmat(obj.valid) = probs_log;
% --- Identify the current true transition
[to_idx, from_idx] = find(obj.valid);
cur_idx = find(dirac==1);
cur_to = to_idx(cur_idx);
cur_from = from_idx(cur_idx);
figure(11); clf;
imagesc(probs_logmat); axis xy; colorbar;
xlabel('From state'); ylabel('To state'); set(gca,'FontSize',10);
hold on;
plot(cur_from, cur_to, 'rs', 'MarkerSize', 10, 'LineWidth', 2, 'MarkerFaceColor', 'none');
hold off;
end
if 1 %training
% dmp is large for the correct transition -> update emphasizes that branch
dL_Dw = dmp .* (yk); % ∂CE/∂(w) - formula (10)
% only start with updates when we are inside the signal
if k > obj.L
obj.w = obj.w - mu * dL_Dw; % (Nf+1)×nFeasible
end
end
% --- Compare-Select (matrix form, min of costs)
v_tilde_mat = inf(obj.nStates, obj.nStates);
v_tilde_mat(obj.valid) = v_tilde;
[pm_next, pred(k,:)] = min(v_tilde_mat, [], 2);
% re-center to keep metrics bounded (decision-invariant)
pm_next = pm_next - min(pm_next);
pm = pm_next;
pm_sto(:,symbol) = pm;
end
% --- Traceback (full; you can window with traceback_depth if desired)
[~, s_end] = min(pm);
viterbi_path = zeros(symbol,1,'uint32');
viterbi_path(symbol) = s_end;
for n = symbol:-1:2
viterbi_path(n-1) = pred(n, viterbi_path(n));
end
y_vit = obj.first_sym(viterbi_path);
y = obj.first_sym(viterbi_path);
if 1 %debug || training
err = sum(y ~= d(1:length(y)));
ser = err./length(y);
fprintf('Epoch: %d - SER: %.1e \n',epoch, ser);
figure(10);
subplot(2,2,1:2);
heatmap(obj.w);
title('Filter')
subplot(2,2,3);
v_tildemat = NaN(obj.nStates, obj.nStates);
v_tildemat(obj.valid) = v_tilde; % log-domain scores
heatmap(v_tildemat);
title('Path Metrics (v_tilde)')
subplot(2,2,4);
scatter(1:symbol,pm_sto,1,'.')
% plot(1:symbol,pm_sto,'LineStyle','none')
title('Path Metric Winners')
end
end
end
end
methods (Access=private)
function k = seq_key(obj, seq)
% Build a stable key string for a sequence row vector in the *same order as combs rows* ([x_k, x_{k-1}, ...])
% Use rounding via sprintf to avoid floating-point issues.
% seq must be a row vector.
k = sprintf(obj.key_fmt, seq);
end
end
end