865 lines
34 KiB
Matlab
865 lines
34 KiB
Matlab
% 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
|
||
|
||
|
||
%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%
|
||
|
||
classdef ML_MLSE < handle
|
||
% ---------------------------------------------------------------------
|
||
% W. Lanneer and Y. Lefevre,
|
||
% “Machine Learning-Based Pre-Equalizers for Maximum Likelihood
|
||
% Sequence Estimation in High-Speed PONs,” EUSIPCO 2023
|
||
% ---------------------------------------------------------------------
|
||
% This implementation reproduces the closed-loop ML-based
|
||
% pre-equalizer training for MLSE, supporting both training and
|
||
% detection (decision-directed) modes.
|
||
% ---------------------------------------------------------------------
|
||
|
||
properties
|
||
sps
|
||
order
|
||
e
|
||
e_tr
|
||
error
|
||
|
||
len_tr
|
||
mu_tr
|
||
epochs_tr
|
||
|
||
dd_mode
|
||
mu_dd
|
||
epochs_dd
|
||
adaptive_mu
|
||
|
||
constellation
|
||
L
|
||
alpha
|
||
DIR
|
||
DIR_flip
|
||
trellis_states
|
||
traceback_depth
|
||
delta
|
||
|
||
% Internal variables
|
||
S
|
||
Nf
|
||
nStates
|
||
nFeasible
|
||
combs
|
||
first_sym
|
||
last_sym
|
||
valid
|
||
valid_to_idx
|
||
valid_from_idx
|
||
w
|
||
|
||
% Fast lookup
|
||
nSym
|
||
key_table
|
||
trans_index
|
||
true_to_state_idx
|
||
|
||
% Debug metrics
|
||
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.001;
|
||
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
|
||
|
||
% ==============================================================
|
||
% PROCESS
|
||
% ==============================================================
|
||
function [X,X_viterbi] = process(obj, X, D)
|
||
% Normalize input RMS
|
||
X = X.normalize("mode","rms");
|
||
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
|
||
|
||
% --- Parameters
|
||
obj.S = obj.nSym;
|
||
obj.Nf = obj.order * obj.sps;
|
||
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{:}).');
|
||
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);
|
||
|
||
% --- Initialize weights
|
||
if isempty(obj.w) || any(size(obj.w) ~= [obj.Nf+1,obj.nFeasible])
|
||
obj.w = randn(obj.Nf+1,obj.nFeasible);
|
||
end
|
||
|
||
% --- Fast lookup tables
|
||
[~, sym_idx_mat] = ismember(obj.combs, obj.constellation);
|
||
key_vals = 1 + sum((sym_idx_mat - 1) .* (obj.nSym .^ (0:obj.L-1)), 2);
|
||
max_key = obj.nSym^obj.L;
|
||
obj.key_table = zeros(max_key,1,'uint32');
|
||
obj.key_table(key_vals) = 1:obj.nStates;
|
||
|
||
obj.trans_index = sparse(obj.nStates,obj.nStates);
|
||
for i = 1:length(obj.valid_from_idx)
|
||
f = obj.valid_from_idx(i);
|
||
t = obj.valid_to_idx(i);
|
||
obj.trans_index(t,f) = i;
|
||
end
|
||
|
||
% ==============================================================
|
||
% TRAINING
|
||
% ==============================================================
|
||
fprintf('\n--- Training mode ---\n');
|
||
obj.equalize(X.signal, D.signal, obj.mu_tr, obj.epochs_tr, obj.len_tr, true);
|
||
obj.e_tr = obj.e;
|
||
|
||
% ==============================================================
|
||
% DECISION-DIRECTED / TESTING
|
||
% ==============================================================
|
||
fprintf('--- Decision-directed / detection mode ---\n');
|
||
[y, y_vit] = obj.equalize(X.signal, D.signal, obj.mu_dd, obj.epochs_dd, X.length, false);
|
||
|
||
X_viterbi = X;
|
||
X.signal = y;
|
||
X_viterbi.signal = y_vit;
|
||
end
|
||
|
||
% ==============================================================
|
||
% EQUALIZE
|
||
% ==============================================================
|
||
function [y,y_ref] = equalize(obj,x,d,mu,epochs,N,training)
|
||
debug = 1;
|
||
showPlots = 1;
|
||
y = zeros(N,1);
|
||
nSymbols = ceil(N/obj.sps);
|
||
|
||
for epoch = 1:epochs
|
||
pm = zeros(obj.nStates,1);
|
||
pred = zeros(nSymbols,obj.nStates,'uint32');
|
||
pm_sto = nan(obj.nStates,nSymbols,'like',pm);
|
||
CE_accum = 0;
|
||
|
||
start_sample = 1;
|
||
end_sample = N;
|
||
start_symbol = 1 + floor((start_sample - 1)/obj.sps);
|
||
|
||
% --- initialize true state
|
||
if numel(d) >= obj.L && start_symbol >= obj.L
|
||
init_seq = d(start_symbol-obj.L+1:start_symbol);
|
||
key_init = obj.seq2key(init_seq);
|
||
true_to_state_idx = obj.key_table(key_init);
|
||
if true_to_state_idx==0, true_to_state_idx=1; end
|
||
else
|
||
true_to_state_idx = uint32(1);
|
||
end
|
||
|
||
for sample = start_sample:obj.sps:end_sample
|
||
symbol = (sample - start_sample)/obj.sps + 1;
|
||
sym_idx = start_symbol + (symbol - 1);
|
||
|
||
% --- Observation window (with delta)
|
||
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)];
|
||
yk = [yk;1];
|
||
|
||
% --- Branch metrics
|
||
c_hat = (yk.' * obj.w).';
|
||
pm = pm - min(pm);
|
||
v_tilde = pm(obj.valid_from_idx) + c_hat;
|
||
|
||
% --- allocate once
|
||
if epoch==1 && symbol==1
|
||
obj.true_to_state_idx = ones(ceil(N/obj.sps),1,'uint32');
|
||
end
|
||
|
||
% --- previous "to" becomes "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
|
||
if sym_idx>=obj.L
|
||
key_to = obj.seq2key(d(sym_idx-obj.L+1:sym_idx));
|
||
state_idx = obj.key_table(key_to);
|
||
if state_idx==0
|
||
state_idx = true_from_state_idx;
|
||
end
|
||
obj.true_to_state_idx(symbol) = state_idx;
|
||
else
|
||
obj.true_to_state_idx(symbol) = true_from_state_idx;
|
||
end
|
||
end
|
||
true_to_state_idx = obj.true_to_state_idx(symbol);
|
||
|
||
% --- fast Dirac creation
|
||
dirac = zeros(obj.nFeasible,1);
|
||
trans_idx = obj.trans_index(true_to_state_idx,true_from_state_idx);
|
||
if trans_idx~=0
|
||
dirac(trans_idx)=1;
|
||
end
|
||
|
||
% --- ensure valid (from,to)
|
||
if ~any(dirac)
|
||
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
|
||
end
|
||
|
||
% ===================================================================
|
||
% TRAINING MODE (weight update)
|
||
% ===================================================================
|
||
if training
|
||
% --- Softmax and CE
|
||
v_shift = -(v_tilde - min(v_tilde));
|
||
v_shift = min(v_shift,100);
|
||
expv = exp(v_shift);
|
||
p = expv./(sum(expv)+eps);
|
||
CE_symbol(symbol) = -log(p(dirac==1)+eps);
|
||
|
||
% --- CE smoothing and adaptive μ
|
||
if sym_idx>obj.L
|
||
CE_smooth(symbol)=0.01*CE_symbol(symbol)+0.99*CE_symbol(symbol-1);
|
||
else
|
||
CE_smooth(symbol)=CE_symbol(symbol);
|
||
end
|
||
CE_accum=CE_accum+CE_symbol(symbol);
|
||
|
||
% --- Gradient update
|
||
dmp=(dirac-p)';
|
||
dL_Dw=(yk).*dmp;
|
||
if sym_idx>=obj.L
|
||
if obj.adaptive_mu
|
||
mu_eff=CE_smooth(symbol);
|
||
mu_eff=max(min(mu_eff,0.2),1e-4);
|
||
else
|
||
mu_eff=mu;
|
||
end
|
||
obj.w=obj.w - mu_eff.*dL_Dw;
|
||
end
|
||
end
|
||
|
||
% ===================================================================
|
||
% DECODING MODE (Viterbi only)
|
||
% ===================================================================
|
||
% Compare-Select (always executed)
|
||
vmat=inf(obj.nStates,obj.nStates);
|
||
vmat(obj.valid)=v_tilde;
|
||
[pm_next,pred(symbol,:)]=min(vmat,[],2);
|
||
pm_next=pm_next-min(pm_next);
|
||
pm=pm_next;
|
||
pm_sto(:,symbol)=pm;
|
||
end
|
||
|
||
% --- Traceback
|
||
[~,s_end]=min(pm);
|
||
vpath=zeros(symbol,1,'uint32');
|
||
vpath(symbol)=s_end;
|
||
for n=symbol:-1:2
|
||
vpath(n-1)=pred(n,vpath(n));
|
||
end
|
||
|
||
y_ref=d(start_symbol:end);
|
||
y=obj.first_sym(vpath);
|
||
|
||
% --- BER/CE reporting and plots
|
||
if training
|
||
err=sum(y~=y_ref(1:length(y)));
|
||
ser=err/length(y);
|
||
try
|
||
ref_bits=PAMmapper(obj.S,0).demap(y_ref(1:length(y)));
|
||
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: %.2e\n',epoch,ber);
|
||
obj.ber(epoch)=ber;
|
||
catch
|
||
fprintf('Epoch %d - SER: %.2e\n',epoch,ser);
|
||
obj.ber(epoch)=ser;
|
||
end
|
||
obj.ce(epoch)=CE_accum/symbol;
|
||
|
||
if debug && mod(epoch,10)==1 && showPlots
|
||
figure(10);clf
|
||
subplot(3,2,1:2);
|
||
imagesc(obj.w);axis xy;colorbar;title('Filter W');
|
||
subplot(3,2,3);
|
||
vtilde_mat=NaN(obj.nStates,obj.nStates);
|
||
vtilde_mat(obj.valid)=v_tilde;
|
||
imagesc(vtilde_mat);axis xy;colorbar;title('Path Metrics (v\_tilde)');
|
||
subplot(3,2,4);
|
||
plot(1:symbol,pm_sto);title('Path Metric Evolution');
|
||
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;
|
||
yyaxis left
|
||
scatter(1:length(obj.ce),obj.ce,10,'s','filled');
|
||
ylabel('Cross Entropy');
|
||
yyaxis right
|
||
scatter(1:length(obj.ber),obj.ber,10,'d','filled');
|
||
set(gca,'YScale','log');
|
||
ylabel('BER (log)');
|
||
xlabel('Epoch');grid on;
|
||
title('Convergence');
|
||
drawnow;
|
||
end
|
||
end
|
||
end
|
||
end
|
||
|
||
% ==============================================================
|
||
% Helper: Sequence → key (always scalar)
|
||
% ==============================================================
|
||
function key = seq2key(obj, seq)
|
||
[~, idx] = ismember(flip(seq), obj.constellation);
|
||
pow = (obj.nSym .^ (0:obj.L-1)).';
|
||
key = 1 + sum((idx(:) - 1) .* pow);
|
||
end
|
||
end
|
||
end
|