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

@@ -121,22 +121,22 @@ classdef ML_MLSE < handle
% ==============================================================
% INITIALIZATION (only before final epoch and detection mode)
% ==============================================================
if epoch == epochs && ~training
if epoch == epochs
% --- Parameters
S = numel(unique(d)); % alphabet size
L = obj.L; % MLSE memory
Nf = L; % filter length
Delta = ceil(L/2); % delay parameter
nStates = S^L;
nFeasible = S^(L-1)*S;
Nf = 9; % filter length
Delta = ceil(Nf/2); % delay parameter
nStates = S^obj.L;
nFeasible = nStates*S;
% --- Trellis mapping
obj.DIR = arburg(y-d, L);
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, L, 1);
pre_comb_cell = mat2cell(pre_comb_mat, ones(1,L), size(pre_comb_mat,2));
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);
@@ -154,7 +154,7 @@ classdef ML_MLSE < handle
[valid_to, valid_from] = find(valid);
% --- 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);
inv2s2 = 1/(2*sigma2);
@@ -166,6 +166,9 @@ classdef ML_MLSE < handle
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
% ==============================================================
@@ -190,7 +193,7 @@ classdef ML_MLSE < handle
obj.e = obj.e + mu * (err * U);
% --- Whitening + MLSE in last epoch
if epoch == epochs && ~training
if epoch == epochs
[y_white(symbol), zi] = filter(obj.DIR,1,y(symbol), zi);
k = symbol;
@@ -210,17 +213,42 @@ classdef ML_MLSE < handle
% --- Extended path metrics
v_tilde = pm(valid_from) + v_hat; % [nFeasible×1]
% --- Compute branch metrics (distance)
bm_vec = -(y_white(k) - v_hat).^2 * inv2s2; % 1×nFeasible
% ===== Gradient update (Algorithm 1) =====
% if training
% --- Survivor selection (vector aggregation)
pm_new_vec = pm(valid_from) + bm_vec.'; % nFeasible×1
pm_next = -inf(nStates,1);
surv_idx = zeros(nStates,1);
for t = 1:nStates
mask = (valid_to==t);
[pm_next(t), arg] = max(pm_new_vec(mask));
surv_idx(t) = valid_from(find(mask,1,'first')-1+arg);
% 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;