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 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.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 % Training Mode n = obj.len_tr; training = 1; obj.equalize(X.signal, D.signal,obj.mu_tr,obj.epochs_dd,n,training); obj.e_tr = obj.e; % 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 % ============================================================== % --- Input padding and preallocation x = [zeros(floor(obj.order/2),1); x; zeros(obj.order,1)]; N_ = N / obj.sps; y = zeros(N_,1); y_white = zeros(N_,1); for epoch = 1:epochs % ============================================================== % INITIALIZATION (only before final epoch and detection mode) % ============================================================== if epoch == epochs % --- Parameters S = numel(unique(d)); % alphabet size Nf = 9; % filter length Delta = ceil(Nf/2); % delay parameter nStates = S^obj.L; nFeasible = nStates*S; % --- Trellis mapping 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, 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); nStates = size(combs,1); % --- Valid transitions valid = false(nStates); for from = 1:nStates for to = 1:nStates if all(combs(to,2:end) == combs(from,1:end-1)) valid(to,from) = true; end end end [valid_to, valid_from] = find(valid); % --- Noise estimation y_ideal = conv(d(1:N_), obj.DIR(:), "same"); sigma2 = mean(abs(y - y_ideal).^2); inv2s2 = 1/(2*sigma2); % --- Allocate vectors and weights pm = zeros(nStates,1); w = zeros(Nf,nFeasible); % filter weights per transition b = zeros(1,nFeasible); % bias terms v_hat = zeros(1,nFeasible); 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 % ============================================================== % RUNTIME LOOP % ============================================================== symbol = 0; for sample = 1:obj.sps:N symbol = symbol + 1; % --- FFE output U = x(obj.order+sample-1:-1:sample); y(symbol,1) = obj.e.' * U; % --- Decision / FFE adaptation if training d_hat = d(symbol); else [~,idx] = min(abs(y(symbol) - obj.constellation)); d_hat = obj.constellation(idx); end err = d_hat - y(symbol); obj.e = obj.e + mu * (err * U); % --- Whitening + MLSE in last epoch if epoch == epochs [y_white(symbol), zi] = filter(obj.DIR,1,y(symbol), zi); k = symbol; % --- Build Δ-delayed observation window y_k i1 = k - Nf + 1 + Delta; i2 = k + Delta; buf = y_white(max(1,i1):min(length(y_white),i2)); padL = max(0,1 - i1); padR = max(0,i2 - length(y_white)); yk = [zeros(padL,1); buf(:); zeros(padR,1)]; % Nf×1 % --- Predict branch metrics for all feasible transitions v_hat = (yk.' * w) + b; % [1×nFeasible] v_hat = v_hat.'; % [nFeasible×1] % --- Extended path metrics v_tilde = pm(valid_from) + v_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 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; % --- Traceback if mod(symbol,obj.traceback_depth) == 0 [~,viterbi_path(symbol)] = max(pm); for n = symbol:-1:symbol-obj.traceback_depth+2 viterbi_path(n-1) = surv_idx(viterbi_path(n)); end end end end % --- Output reconstructed path if epoch == epochs && ~training y_vit = first_sym(viterbi_path); end end end end end