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 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"); obj.constellation = unique(D.signal); 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(unique(D.signal)); % 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(unique(D.signal),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{:}).'); 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 all(size(obj.w) ~= [obj.Nf+1,obj.nFeasible]) obj.w = zeros(obj.Nf+1,obj.nFeasible); % filter weights per transition + bias tap end % obj.w = randn(obj.Nf+1,obj.nFeasible); % ============================================================== % 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); for epoch = 1:epochs pm = zeros(obj.nStates,1); c_hat = zeros(1,obj.nFeasible); v_tilde = zeros(1,obj.nFeasible); pred = zeros(N, obj.nStates, 'uint32'); % ============================================================== % RUNTIME LOOP % ============================================================== 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 = (yk.' * obj.w); % [1×nFeasible] c_hat = c_hat.'; % [nFeasible×1] % --- Extended path metrics v_tilde = pm(obj.valid_from_idx) + c_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 true_from_state_idx = find(ismember(obj.combs, flip(prev_seq).', 'rows'));% find state indices in trellis %if ~(true_from_state_idx == true_to_state_idx), warning('Impossible state transition?!'), end curr_seq = d(k-obj.L+1:k); % next state symbols true_to_state_idx = find(ismember(obj.combs, flip(curr_seq).', 'rows'));% find state indices in trellis else % not enough history yet true_from_state_idx = 1; true_to_state_idx = 1; 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 % true_from and true_to are (1) -> (-1) % thus dirac = [1,0,0,0]' 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. if sum(dirac) == 0 warning('whats happening?!'); end % softmax over -v_tilde, one-hot target t % 4 feasible transitions: [0,0][0,1][1,0][1,1] #not in % order here! % first round: % v_tilde is zero % -> exp(-0) = 1 % -> 1/(1+1+1+1) % -> 1/4 % -> each transition is equally likely? % p = exp(-v_tilde); p = exp(-(v_tilde-max(v_tilde))); p = p./sum(p); % found in formula (9) and (19) % 1-0.25 = 0.75 % 0-0.25 = -0.25 % what happens here at the zero? Is this some log % probability - larger zero is likely; lower zero is % unlikely? But this is based on known information (training) dmp = (dirac - p)'; if mod(symbol,128) == 1 && debug % --- Normalize and compute probabilities v_norm = v_tilde - max(v_tilde); % numerical stability 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_mat(obj.valid) = probs_lin; % linear-space probabilities probs_logmat = nan(obj.nStates, obj.nStates); probs_logmat(obj.valid) = probs_log; % log-domain scores % --- 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); % --- Plot using imagesc (supports hold) figure(11); clf; imagesc(probs_logmat); axis xy; % origin top-left % colormap(parula); colorbar; xlabel('From state'); ylabel('To state'); set(gca,'FontSize',10); % --- Overlay current transition hold on; plot(cur_from, cur_to, 'rs', ... 'MarkerSize', 10, 'LineWidth', 2, 'MarkerFaceColor', 'none'); hold off; end if 1 %training % the correct state gets a high update, we weight this % with the input signal, from here the signals are not % "understandable" % dmp is large for the correct transition to update % only this one! dL_Dw = dmp .* (yk); % ∂CE/∂(w) - formula (10) % from paper: We have observed in simulations that ignoring the derivative of vk−1(s′) during training yields negligible loss after convergence. % my note: so the update direction is the same for b and w? %only start with updates when we are inside the signal if k > obj.L % actual filter training updates: obj.w = obj.w - mu * dL_Dw;% Nf×nFeasible end end % compare select v_tilde_mat = inf(obj.nStates, obj.nStates); v_tilde_mat(obj.valid) = v_tilde; [pm, pred(k,:)] = min(v_tilde_mat, [], 2); pm_sto(:,symbol) = pm; end [~, s_end] = max(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:N,pm_sto,1,'.') plot(1:symbol,pm_sto) title('Path Metric Winners') end end end end end