classdef ML_MLSE < handle % ALGORITHM DESCRIBED IN: % W. Lanneer and Y. Lefevre, “Machine Learning-Based Pre-Equalizers for % Maximum Likelihood Sequence Estimation in High-Speed PONs,” % in 2023 31st European Signal Processing Conference % Further ML Refs: % https://machinelearningmastery.com/cross-entropy-for-machine-learning/ % https://docs.pytorch.org/docs/stable/generated/torch.nn.CrossEntropyLoss.html % The central idea is to overcome the (white-) noise assumption within the previously described % Viterbi algorithm, more precisely a closed-loop optimization is proposed that finds a suitable % filter-set to directly compute the branch metrics c_k (s,s^' ). These can directly be used to % carry out the conventional Viterbi algorithm. The system consists of S^L S=F linear FIR filters, % combined with one bias coefficient respectively. These filters take the received input samples to % compute the branch metrics estimates (c_k ) ̂(s,s^' ) according toThe central idea is to overcome % the (white-) noise assumption within the previously described Viterbi algorithm, more precisely % a closed-loop optimization is proposed that finds a suitable filter-set to directly compute the % branch metrics c_k (s,s^' ). These can directly be used to carry out the conventional Viterbi % algorithm. The system consists of S^L S=F linear FIR filters, combined with one bias coefficient % respectively. These filters take the received input samples to compute the branch metrics % estimates. Finally, the usual Viterbi is carried out... % Recommended Settings and some findings: % Requires many training epochs. According to ML people, 100,200 or % even up to 1000 epochs are normal for ML-convergence % The mu parameter _can_ be adaptive - using the cross entropy and when % analyzing the isolated training it looks very promisig. However, is % later use I found this is not as stable as a fixed learning rate. % mu = 0.1 worked good for me % Longer orders/ filter length are not always better. For me order=11 % was good. % Delay factor (delta) is good when the order is also increased. With % order = 11, a delta of =4 shows good results 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 adaptive_mu 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 % --- New: fast state lookup --- true_to_state_idx state_dict % containers.Map: key(sequence)->state index key_fmt = '%.8g_'; % key format for sequence strings nSym % |constellation| ber = [] ce = ones(1,1); 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.adaptive_mu = 1; 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"); % Use sorted constellation for deterministic mapping obj.constellation = sort(unique(D.signal),'ascend'); obj.nSym = numel(obj.constellation); 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(obj.constellation); % 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(obj.constellation,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{:}).'); % rows: states, columns: [x_k, x_{k-1}, ...] 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 isempty(obj.w) || any(size(obj.w) ~= [obj.Nf+1,obj.nFeasible]) obj.w = zeros(obj.Nf+1,obj.nFeasible); % filter weights per transition + bias tap obj.w = randn(obj.Nf+1,obj.nFeasible); end % --- Precompute dictionary for fast state lookup (sequence -> state) keys = cell(obj.nStates,1); for i = 1:obj.nStates keys{i} = obj.seq_key(obj.combs(i,:)); % combs row is already [x_k, x_{k-1}, ...] end obj.state_dict = containers.Map(keys, 1:obj.nStates); % ============================================================== % 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_ref] = equalize(obj,x,d,mu,epochs,N,training) % ============================================================== % FFE + Whitening + ML-Based Branch Metric Estimation + Viterbi % ============================================================== debug = 1; showPlots = 1; % --- Input padding and preallocation y = zeros(N,1); % number of symbol steps in this block nSymbols = ceil(N/obj.sps); for epoch = 1:epochs % state metrics (log-domain costs): keep as column [nStates×1] pm = zeros(obj.nStates,1); % v_{k-1}(s′) c_hat = zeros(1,obj.nFeasible); v_tilde = zeros(1,obj.nFeasible); pred = zeros(nSymbols, obj.nStates, 'uint32'); pm_sto = nan(obj.nStates, nSymbols,'like',pm); CE_accum = 0; %%% START IDX if training max_start = length(x) - ( (ceil(N/obj.sps)-1)*obj.sps + 1 ); max_start = max(1, max_start); % safety start_sample = randi([1, max_start], 1); %rnd training; not really good start_sample = 1; end_sample = start_sample + (ceil(N/obj.sps)-1)*obj.sps; else start_sample = 1;%obj.len_tr; end_sample = N; end start_symbol = 1 + floor((start_sample - 1)/obj.sps); % ABSOLUTE symbol index if numel(d) >= obj.L && start_symbol >= obj.L init_seq = d(start_symbol-obj.L+1 : start_symbol); % [d_k-L+1 ... d_k] true_to_state_idx = obj.state_dict(obj.seq_key(flip(init_seq))); % [d_k ... d_k-L+1] else % Not enough history – fall back to state 1 true_to_state_idx = uint32(1); end symbol = 0; for sample = start_sample:obj.sps:end_sample symbol = symbol + 1; k = symbol; sym_idx = start_symbol + (symbol - 1); % --- 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 c_hat = (yk.' * obj.w); % [1×nFeasible] c_hat = c_hat.'; % [nFeasible×1] % --- Extended path metrics: v_tilde = pm(from) + c_hat % normalize pm to avoid growth (invariant to additive const) pm = pm - min(pm); v_tilde = pm(obj.valid_from_idx) + c_hat; % [nFeasible×1] % ===== Gradient update (Algorithm 1) ===== if 1 %training % --- allocate storage once if epoch == 1 && symbol == 1 obj.true_to_state_idx = ones(ceil(N/obj.sps),1,'uint32'); end % --- previous "to" becomes current "from" if symbol > 1 true_from_state_idx = obj.true_to_state_idx(symbol-1); else true_from_state_idx = 1; end % --- compute or reuse "to" state if epoch == 1 % only compute in first epoch if sym_idx >= obj.L key_to = obj.seq_key(flip(d(sym_idx-obj.L+1 : sym_idx))); if isKey(obj.state_dict, key_to) obj.true_to_state_idx(symbol) = obj.state_dict(key_to); else obj.true_to_state_idx(symbol) = true_from_state_idx; end else obj.true_to_state_idx(symbol) = true_from_state_idx; end end % --- reuse cached state from second epoch onward true_to_state_idx = obj.true_to_state_idx(symbol); % --- ensure valid (from,to) dirac = zeros(obj.nFeasible,1); mask = obj.valid_from_idx==true_from_state_idx & ... obj.valid_to_idx ==true_to_state_idx; if any(mask) dirac(mask) = 1; else idx = find(obj.valid_from_idx==true_from_state_idx,1,'first'); dirac(idx) = 1; obj.true_to_state_idx(symbol) = obj.valid_to_idx(idx); end % softmax over -v_tilde (numerically safe shift) v_shift = -(v_tilde - min(v_tilde)); % shift to small positive numbers v_shift = min(v_shift, 100); % clamp exponent argument (≈ exp(50)=3e21) expv = exp(v_shift); p = expv ./ (sum(expv) + eps); % for logging only: CE_symbol(symbol) = -log(p(dirac==1) + eps); if sym_idx > obj.L CE_smooth(symbol) = 0.01*CE_symbol(symbol) + 0.99*CE_smooth(symbol-1); else if epoch > 1 CE_smooth(symbol) = obj.ce(end); %use ce from last epoch or =1 for very first round?! else CE_smooth(symbol) = CE_symbol(symbol); end end CE_accum = CE_symbol(symbol) + CE_accum; % gradient term (t - p) dmp = (dirac - p)'; % 1×nFeasible % Per-feature gradient; implicit expansion gives (Nf+1)×nFeasible dL_Dw = (yk) .* dmp; % Start updates only when the ABSOLUTE symbol index has ≥ L history if sym_idx >= obj.L if obj.adaptive_mu mu_eff = CE_smooth(sym_idx); mu_eff = max(min(mu_eff, 0.2), 1e-4); else mu_eff = mu; end obj.w = obj.w - mu_eff .* dL_Dw; % (Nf+1)×nFeasible end % if debug && epoch > 2 % figure(100); % subplot(4,1,1); % heatmap(p'); % title('Probs') % subplot(4,1,2); % heatmap(dmp); % title('Update') % subplot(4,1,3); % heatmap(dL_Dw); % title('Update') % subplot(4,1,4); % heatmap(bj.w); % title('Update') % % end end % --- Compare-Select (matrix form, min of costs) v_tilde_mat = inf(obj.nStates, obj.nStates); v_tilde_mat(obj.valid) = v_tilde; [pm_next, pred(k,:)] = min(v_tilde_mat, [], 2); % re-center to keep metrics bounded (decision-invariant) pm_next = pm_next - min(pm_next); pm = pm_next; pm_sto(:,symbol) = pm; end % --- Traceback (full; you can window with traceback_depth if desired) [~, s_end] = min(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_ref = d(start_symbol:end); y = obj.first_sym(viterbi_path); if debug && training sym_start = start_symbol; sym_end = start_symbol + symbol - 1; ref_slice = d(sym_start : sym_end); err = sum(y ~= ref_slice(1:numel(y))); try ref_bits = PAMmapper(obj.S,0).demap(ref_slice); eq_bits = PAMmapper(obj.S,0).demap(y); [~, ~, ber, ~] = calc_ber(ref_bits, eq_bits, "skip_front", 10, "skip_end", 10, "returnErrorLocation", 1); fprintf('Epoch: %d - BER: %.1e \n',epoch, ber); obj.ber(epoch) = ber; catch ser = err./length(y); fprintf('Epoch: %d - SER: %.1e \n',epoch, ser); end obj.ce(epoch) = CE_accum./symbol; if showPlots figure(10);clf subplot(3,2,1:2); heatmap(obj.w); title('Filter') subplot(3,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(3,2,4); scatter(1:symbol,pm_sto,1,'.') title('Path Metric Winners') subplot(3,2,5);hold on scatter(1:symbol,CE_symbol,1,'.'); scatter(1:symbol,CE_smooth,1,'.') title('Cross Entropy') subplot(3,2,6); hold on % Left y-axis: Cross Entropy (linear) yyaxis left scatter(1:length(obj.ce), obj.ce, 10, 's', 'filled') ylabel('Cross Entropy') % Right y-axis: BER (logarithmic) yyaxis right scatter(1:length(obj.ber), obj.ber, 10, 'd', 'filled') set(gca, 'YScale', 'log') ylabel('BER (log scale)') xlim([1, epochs]) xlabel('Epoch') title('Cross Entropy // BER') grid on drawnow end end end end end methods (Access=private) function k = seq_key(obj, seq) % Build a stable key string for a sequence row vector in the *same order as combs rows* ([x_k, x_{k-1}, ...]) % Use rounding via sprintf to avoid floating-point issues. % seq must be a row vector. k = sprintf(obj.key_fmt, seq); end end end