diff --git a/Classes/04_DSP/Equalizer/FFE_MLSE.m b/Classes/04_DSP/Equalizer/FFE_MLSE.m index e7a4bbd..d94da34 100644 --- a/Classes/04_DSP/Equalizer/FFE_MLSE.m +++ b/Classes/04_DSP/Equalizer/FFE_MLSE.m @@ -1,4 +1,4 @@ -classdef ML_MLSE < handle +classdef FFE_MLSE < handle % Implementation of plain and simple FFE. % 1) Training mode (stable performance when you use NLMS) % 2) Decision directed mode @@ -37,7 +37,7 @@ classdef ML_MLSE < handle end methods - function obj = ML_MLSE(options) + function obj = FFE_MLSE(options) arguments(Input) options.sps = 2; diff --git a/Classes/04_DSP/Equalizer/ML_MLSE.m b/Classes/04_DSP/Equalizer/ML_MLSE.m index db6c686..65eefb8 100644 --- a/Classes/04_DSP/Equalizer/ML_MLSE.m +++ b/Classes/04_DSP/Equalizer/ML_MLSE.m @@ -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; diff --git a/projects/ML_based_MLSE/model.m b/projects/ML_based_MLSE/model.m index f74e5fe..d1da091 100644 --- a/projects/ML_based_MLSE/model.m +++ b/projects/ML_based_MLSE/model.m @@ -99,15 +99,15 @@ for r = 1:length(fsym) %% Implement DSP directly here: mu_lms = 0.0005; + % tic + % eq = FFE_MLSE("epochs_tr",5,"epochs_dd",5,"len_tr",2^13,"mu_dd",mu_lms,"mu_tr",mu_lms,"order",50,"sps",2,"traceback_depth",32,"L",1); + % [Eq_signal,Vit_signal] = eq.process(Rx_sig_2sps,Symbols); + % toc + tic eq = ML_MLSE("epochs_tr",5,"epochs_dd",5,"len_tr",2^13,"mu_dd",mu_lms,"mu_tr",mu_lms,"order",50,"sps",2,"traceback_depth",32,"L",1); [Eq_signal,Vit_signal] = eq.process(Rx_sig_2sps,Symbols); toc - - % tic - % eq = FFE_MLSE("epochs_tr",5,"epochs_dd",5,"len_tr",2^13,"mu_dd",mu_lms,"mu_tr",mu_lms,"order",50,"sps",2,"traceback_depth",32,"L",3); - % [Eq_signal,Vit_signal] = eq.process(Rx_sig_2sps,Symbols); - % toc Eq_bits = PAMmapper(M, 0, "eth_style", 0).demap(Eq_signal); [~, errors, ber, ~] = calc_ber(Eq_bits.signal, Tx_bits.signal, "skip_front", 0, "skip_end", 0, "returnErrorLocation", 1); @@ -122,12 +122,11 @@ for r = 1:length(fsym) - %% optimize smth. tr_len = 2.^[2:15]; tr_len = floor(tr_len); ber = zeros(size(tr_len)); - parfor m = 1:numel(tr_len) + for m = 1:numel(tr_len) mu_lms = 0.0005; eq = ML_MLSE("epochs_tr",5,"epochs_dd",5,"len_tr",2^13,"mu_dd",mu_lms,"mu_tr",mu_lms,"order",50,"sps",2,"traceback_depth",tr_len(m)); @@ -150,19 +149,6 @@ for r = 1:length(fsym) set(gca,'YScale','log'); - - - - - - - - - - - - - %% RUN Comparison len_tr = 4096*2;