364 lines
14 KiB
Matlab
364 lines
14 KiB
Matlab
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∗ k−1) = 1 if the extended state (s, s′) corresponds to the true realized states (s∗ k, s∗ k−1), 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
|