Developing the ML-enhanced MLSE :-)
This commit is contained in:
@@ -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
|
||||
% =========================================
|
||||
|
||||
% compare–select 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;
|
||||
|
||||
|
||||
Reference in New Issue
Block a user