Files
imdd_silas/Classes/04_DSP/Equalizer/ML_MLSE.m

865 lines
34 KiB
Matlab
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
% 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