Files
imdd_silas/Classes/04_DSP/Equalizer/ML_MLSE.m
Silas Oettinghaus ac95aef1e0 Works at 1 SPS
NWhitened Signal or FFE Output are okay.
mu is approx. 0.15
delta is 0
order is 4 (higher is not better yet)
2025-10-30 14:29:14 +01:00

365 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
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");
obj.constellation = unique(D.signal);
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(unique(D.signal)); % alphabet size
obj.Nf = obj.order; % 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(unique(D.signal),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{:}).');
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 all(size(obj.w) ~= [obj.Nf+1,obj.nFeasible])
obj.w = zeros(obj.Nf+1,obj.nFeasible); % filter weights per transition + bias tap
end
% obj.w = randn(obj.Nf+1,obj.nFeasible);
% ==============================================================
% 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);
for epoch = 1:epochs
pm = zeros(obj.nStates,1);
c_hat = zeros(1,obj.nFeasible);
v_tilde = zeros(1,obj.nFeasible);
pred = zeros(N, obj.nStates, 'uint32');
% ==============================================================
% RUNTIME LOOP
% ==============================================================
symbol = 0;
for sample = 1:obj.sps:N
symbol = symbol + 1;
k = symbol;
% --- Build Δ-delayed observation window y_k
i1 = k - obj.Nf + 1 + obj.delta;
i2 = k + 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 = (yk.' * obj.w); % [1×nFeasible]
c_hat = c_hat.'; % [nFeasible×1]
% --- Extended path metrics
v_tilde = pm(obj.valid_from_idx) + c_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
true_from_state_idx = find(ismember(obj.combs, flip(prev_seq).', 'rows'));% find state indices in trellis
%if ~(true_from_state_idx == true_to_state_idx), warning('Impossible state transition?!'), end
curr_seq = d(k-obj.L+1:k); % next state symbols
true_to_state_idx = find(ismember(obj.combs, flip(curr_seq).', 'rows'));% find state indices in trellis
else
% not enough history yet
true_from_state_idx = 1;
true_to_state_idx = 1;
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
% true_from and true_to are (1) -> (-1)
% thus dirac = [1,0,0,0]'
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.
if sum(dirac) == 0
warning('whats happening?!');
end
% softmax over -v_tilde, one-hot target t
% 4 feasible transitions: [0,0][0,1][1,0][1,1] #not in
% order here!
% first round:
% v_tilde is zero
% -> exp(-0) = 1
% -> 1/(1+1+1+1)
% -> 1/4
% -> each transition is equally likely?
% p = exp(-v_tilde);
p = exp(-(v_tilde-max(v_tilde)));
p = p./sum(p); % found in formula (9) and (19)
% 1-0.25 = 0.75
% 0-0.25 = -0.25
% what happens here at the zero? Is this some log
% probability - larger zero is likely; lower zero is
% unlikely? But this is based on known information (training)
dmp = (dirac - p)';
if mod(symbol,128) == 1 && debug
% --- Normalize and compute probabilities
v_norm = v_tilde - max(v_tilde); % numerical stability
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_mat(obj.valid) = probs_lin; % linear-space probabilities
probs_logmat = nan(obj.nStates, obj.nStates);
probs_logmat(obj.valid) = probs_log; % log-domain scores
% --- 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);
% --- Plot using imagesc (supports hold)
figure(11); clf;
imagesc(probs_logmat);
axis xy; % origin top-left
% colormap(parula);
colorbar;
xlabel('From state');
ylabel('To state');
set(gca,'FontSize',10);
% --- Overlay current transition
hold on;
plot(cur_from, cur_to, 'rs', ...
'MarkerSize', 10, 'LineWidth', 2, 'MarkerFaceColor', 'none');
hold off;
end
if 1 %training
% the correct state gets a high update, we weight this
% with the input signal, from here the signals are not
% "understandable"
% dmp is large for the correct transition to update
% only this one!
dL_Dw = dmp .* (yk); % ∂CE/∂(w) - formula (10)
% from paper: We have observed in simulations that ignoring the derivative of vk1(s) during training yields negligible loss after convergence.
% my note: so the update direction is the same for b and w?
%only start with updates when we are inside the signal
if k > obj.L
% actual filter training updates:
obj.w = obj.w - mu * dL_Dw;% Nf×nFeasible
end
end
% compare select
v_tilde_mat = inf(obj.nStates, obj.nStates);
v_tilde_mat(obj.valid) = v_tilde;
[pm, pred(k,:)] = min(v_tilde_mat, [], 2);
pm_sto(:,symbol) = pm;
end
[~, s_end] = max(pm);
viterbi_path = zeros(N,1,'uint32');
viterbi_path(N) = s_end;
for n = N:-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:N,pm_sto,1,'.')
plot(1:N,pm_sto)
title('Path Metric Winners')
end
end
end
end
end