From 39bc8243fc0b828529638dcad89e37e215c4d367 Mon Sep 17 00:00:00 2001 From: "silas (home)" Date: Mon, 3 Nov 2025 08:17:39 +0100 Subject: [PATCH] Few more scripts to evaluate new ML-MLSE Equalizer - which is not better!!! :-( --- Classes/04_DSP/Equalizer/ML_MLSE.m | 182 +++++++++------- projects/IMDD_base_system/imdd_it.m | 2 +- projects/ML_based_MLSE/minimal_model_gpt.m | 213 ------------------- projects/ML_based_MLSE/model.m | 66 +++--- projects/ML_based_MLSE/rate_evaluation.m | 109 ++++++++++ projects/ML_based_MLSE/rop_evaluation.m | 95 +++++++++ projects/ML_based_MLSE/standard_link_model.m | 95 +++++++++ 7 files changed, 443 insertions(+), 319 deletions(-) delete mode 100644 projects/ML_based_MLSE/minimal_model_gpt.m create mode 100644 projects/ML_based_MLSE/rate_evaluation.m create mode 100644 projects/ML_based_MLSE/rop_evaluation.m create mode 100644 projects/ML_based_MLSE/standard_link_model.m diff --git a/Classes/04_DSP/Equalizer/ML_MLSE.m b/Classes/04_DSP/Equalizer/ML_MLSE.m index e09d2cd..bb327b1 100644 --- a/Classes/04_DSP/Equalizer/ML_MLSE.m +++ b/Classes/04_DSP/Equalizer/ML_MLSE.m @@ -53,6 +53,8 @@ classdef ML_MLSE < handle state_dict % containers.Map: key(sequence)->state index key_fmt = '%.8g_'; % key format for sequence strings nSym % |constellation| + + ber = [] end methods @@ -182,7 +184,7 @@ classdef ML_MLSE < handle % ============================================================== % FFE + Whitening + ML-Based Branch Metric Estimation + Viterbi % ============================================================== - debug = 0; + debug = 1; % --- Input padding and preallocation y = zeros(N,1); @@ -199,19 +201,34 @@ classdef ML_MLSE < handle pred = zeros(nSymbols, obj.nStates, 'uint32'); pm_sto = nan(obj.nStates, nSymbols,'like',pm); - % --- Initialize "true" trellis state for training (shift-register style) - % expect sequences in chronological order [x_{k-L+1}:x_k], but combs rows are [x_k, x_{k-1}, ...] - if numel(d) >= obj.L - init_seq = d(1:obj.L); % [x_1 ... x_L] - true_to_state_idx = obj.state_dict(obj.seq_key(flip(init_seq))); % flip to [x_L, x_{L-1}, ...] + + %%% 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 - true_to_state_idx = 1; + start_sample = 1; + 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 = 1:obj.sps:N + 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; @@ -232,78 +249,72 @@ classdef ML_MLSE < handle v_tilde = pm(obj.valid_from_idx) + c_hat; % [nFeasible×1] % ===== Gradient update (Algorithm 1) ===== - % if training - - if k > obj.L - % shift-register: previous "to" becomes current "from" - true_from_state_idx = true_to_state_idx; - - % current "to" from data window - curr_seq = d(k-obj.L+1:k); - key_to = obj.seq_key(flip(curr_seq)); - if isKey(obj.state_dict, key_to) - true_to_state_idx = obj.state_dict(key_to); - else - % fall back safely (should not happen with proper constellation) - true_to_state_idx = true_from_state_idx; - end - else - % not enough history yet - true_from_state_idx = 1; - true_to_state_idx = true_to_state_idx; % keep init - 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 - - % Dirac delta over correct extended transition (from,to) - 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. - - % softmax over -v_tilde (numerically safe shift) - p = exp(-(v_tilde - max(v_tilde))); - p = p./sum(p); % found in formula (9) and (19) - - % gradient term (t - p) - dmp = (dirac - p)'; % 1×nFeasible - - if mod(symbol,128) == 1 && debug - % --- Normalize and compute probabilities - v_norm = v_tilde - max(v_tilde); - 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_logmat = nan(obj.nStates, obj.nStates); - probs_mat(obj.valid) = probs_lin; - probs_logmat(obj.valid) = probs_log; - - % --- 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); - - figure(11); clf; - imagesc(probs_logmat); axis xy; colorbar; - xlabel('From state'); ylabel('To state'); set(gca,'FontSize',10); - hold on; - plot(cur_from, cur_to, 'rs', 'MarkerSize', 10, 'LineWidth', 2, 'MarkerFaceColor', 'none'); - hold off; - end if 1 %training - % dmp is large for the correct transition -> update emphasizes that branch - dL_Dw = dmp .* (yk); % ∂CE/∂(w) - formula (10) - % only start with updates when we are inside the signal - if k > obj.L - obj.w = obj.w - mu * dL_Dw; % (Nf+1)×nFeasible + % previous "to" becomes current "from" (shift-register) + true_from_state_idx = true_to_state_idx; + + % --- Build current "to" state from ABSOLUTE symbol index + if sym_idx >= obj.L + curr_seq = d(sym_idx-obj.L+1 : sym_idx); % [d_k-L+1 ... d_k] + key_to = obj.seq_key(flip(curr_seq)); % -> [d_k ... d_k-L+1] + if isKey(obj.state_dict, key_to) + true_to_state_idx = obj.state_dict(key_to); + else + % Fall back safely (should not happen with proper constellation) + true_to_state_idx = true_from_state_idx; + end + else + % Not enough history yet for a full L-symbol state + % keep previous 'to' and 'from' + true_to_state_idx = true_to_state_idx; + true_from_state_idx = true_from_state_idx; end + + % Dirac delta over correct extended transition (from,to) + dirac = zeros(obj.nFeasible,1); + dirac(obj.valid_from_idx==true_from_state_idx & ... + obj.valid_to_idx ==true_to_state_idx) = 1; + + % softmax over -v_tilde (numerically safe shift) + p = exp(-(v_tilde - max(v_tilde))); + p = p./(sum(p)+eps); + + % 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 + + + obj.w = obj.w - ones(size(dL_Dw,1),1).*mu .* dL_Dw; % (Nf+1)×nFeasible + % obj.w = obj.w - mu * 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; @@ -327,10 +338,21 @@ classdef ML_MLSE < handle 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); + 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))); + + 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; + + % ser = err./length(y); + % fprintf('Epoch: %d - SER: %.1e \n',epoch, ser); figure(10); subplot(2,2,1:2); @@ -347,6 +369,8 @@ classdef ML_MLSE < handle scatter(1:symbol,pm_sto,1,'.') % plot(1:symbol,pm_sto,'LineStyle','none') title('Path Metric Winners') + + drawnow end end end diff --git a/projects/IMDD_base_system/imdd_it.m b/projects/IMDD_base_system/imdd_it.m index d2fdf3c..bd2461d 100644 --- a/projects/IMDD_base_system/imdd_it.m +++ b/projects/IMDD_base_system/imdd_it.m @@ -15,7 +15,7 @@ if 1 wh.addStorage("ber"); % wh = submit_simulations(wh,"parallel",0,"simulation_mode",0); - wh = submit_handle(@imdd_model,wh,"parallel",1); + wh = submit_handle(@imdd_model,wh,"parallel",0); end diff --git a/projects/ML_based_MLSE/minimal_model_gpt.m b/projects/ML_based_MLSE/minimal_model_gpt.m deleted file mode 100644 index 9536774..0000000 --- a/projects/ML_based_MLSE/minimal_model_gpt.m +++ /dev/null @@ -1,213 +0,0 @@ -function minimal_model_gpt -% Minimal binary IM/DD + ISI (M=3), AWGN. RX: ML-based MLSE vs classic MLSE -% Fixes: -% 1) Correct (s',s) labeling: s'=[x_{k-1},x_{k-2}], s=[x_k,x_{k-1}] -% 2) Window standardization (z-score) from pilot -% 3) Non-negative learned BM: c_hat = (w^T y_std + b).^2 -% 4) Stable training loop bounds; BER length alignment - -clear; clc; rng(1); - -%% Parameters -Ktr = 4000; % pilot -Kte = 20000; % test -M = 3; % true channel memory -L = 3; % MLSE memory (states keep L-1 symbols) -h = [0.55 1.00 0.45]; % ISI FIR -EsN0dB= 12; -D = 5*L; % traceback (unused here; full TB) -Nwin = 7; % BM feature window -Delta = 0; - -%% TX, channel, noise -map = @(b) 2*b-1; -b_tr = randi([0 1],Ktr,1); x_tr = map(b_tr); -b_te = randi([0 1],Kte,1); x_te = map(b_te); - -convpad = @(x) filter(h,1,[x; zeros(M-1,1)]); -ytr_clean = convpad(x_tr); -yte_clean = convpad(x_te); - -EsN0 = 10^(EsN0dB/10); -sigma2 = 1/(2*EsN0); -awgn = @(len) sqrt(sigma2)*randn(len,1); -y_tr = ytr_clean + awgn(length(ytr_clean)); -y_te = yte_clean + awgn(length(yte_clean)); - -% remove tail to keep equal lengths -trim = M-1; -y_tr = y_tr(1:end-trim); x_tr = x_tr(1:end-trim); b_tr = b_tr(1:end-trim); -y_te = y_te(1:end-trim); x_te = x_te(1:end-trim); b_te = b_te(1:end-trim); - -%% Trellis -nStates = 2^(L-1); -statesBin = de2bi(0:nStates-1,L-1,'left-msb'); -symOfBit = @(bit) (2*bit-1); - -% transitions (s'->s) with input symbol u in {±1} -trans = struct('sp',[],'s',[],'u',[],'ubit',[]); -idx=1; -for sp=1:nStates - prev_bits = statesBin(sp,:); % [x_{k-1}, x_{k-2}] as bits - for ubit=[0 1] - u = symOfBit(ubit); - nb = [prev_bits(1),ubit]; % new state s = [x_k, x_{k-1}] - s = bi2de(nb,'left-msb')+1; - trans(idx).sp=sp; trans(idx).s=s; trans(idx).u=u; trans(idx).ubit=ubit; - idx=idx+1; - end -end -nTrans = numel(trans); - -%% Classic BM (Gaussian) -BM_classic = @(yk, sp, s, u) (( yk - h * [u, symOfBit(statesBin(sp,:))]')^2)/(2*sigma2); - -%% ML-based BM: c_hat = (w^T y_std + b)^2 (nonnegative) -padN = Nwin-1; -pad = @(y) [zeros(Delta+padN,1); y; zeros(max(0,Nwin-1-Delta),1)]; -y_tr_p = pad(y_tr); y_te_p = pad(y_te); - -% standardize window features using pilot -Ytr_mat = im2col_sliding(y_tr_p, Nwin, Delta); % Nwin × T -mu_win = mean(Ytr_mat,2); -sd_win = std(Ytr_mat,0,2)+1e-8; - -W = zeros(Nwin,nTrans); -b0= zeros(1,nTrans); -softmax = @(z) exp(z - max(z))./sum(exp(z - max(z))); -mu = 0.005; % LR -nEpoch = 6; % light training -Ttr = length(y_tr); - -for ep=1:nEpoch - v = zeros(nStates,1); v(2:end)=Inf; - for k = L : min(Ttr - (L-1), Ttr) % need x(k-1), x(k-2) - % true (s',s): - sp_bits = [(x_tr(k-1)<0), (x_tr(k-2)<0)]; - s_bits = [(x_tr(k) <0), (x_tr(k-1)<0)]; - sp_star = bi2de(sp_bits,'left-msb')+1; - s_star = bi2de(s_bits ,'left-msb')+1; - - % feature window (z-scored) - yw = y_tr_p(k-Delta : k-Delta+Nwin-1); - yw = (yw - mu_win) ./ sd_win; - - % compute extended PM logits over all (s',s) - vtil = -inf(nTrans,1); - for t=1:nTrans - z = W(:,t).'*yw + b0(t); - c_hat = z*z; % (.)^2 nonnegative BM - vtil(t) = -( v(trans(t).sp) + c_hat); - end - pext = softmax(vtil).'; - t_star = find([trans.sp]==sp_star & [trans.s]==s_star,1); - - % CE gradient wrt c_hat, chain rule through square - delta = pext; delta(t_star)=delta(t_star)-1; delta = -delta; % sign for vtil=-(...) - for t=1:nTrans - z = W(:,t).'*yw + b0(t); - dc_dz = 2*z; % d(z^2)/dz - g = delta(t) * dc_dz; % ∂L/∂z - W(:,t) = W(:,t) - mu * g * yw; - b0(t) = b0(t) - mu * g; - end - - % PM update for next step - v_next = inf(nStates,1); - for t=1:nTrans - sp=trans(t).sp; s=trans(t).s; - z = W(:,t).'*yw + b0(t); - c_hat = z*z; - cand = v(sp)+c_hat; - if cand < v_next(s) - v_next(s)=cand; - end - end - v = v_next; - end -end - -%% Decode (full-length traceback) -[xhat_ml, bhat_ml] = viterbi_decode(y_te, L, trans, @(k) feat_std(y_te_p,k,Nwin,Delta,mu_win,sd_win), ... - @(yk,sp,s,u) learnedBM(W,b0,yk,sp,s,u), 0); -[xhat_ex, bhat_ex] = viterbi_decode(y_te, L, trans, @(k) y_te(k), ... - @(yk,sp,s,u) BM_classic(yk,sp,s,u), 0); - -%% BER (align by min length) -Lmin = min([length(bhat_ml), length(bhat_ex), length(b_te)]); -BER_ml = mean(bhat_ml(1:Lmin) ~= (b_te(1:Lmin)>0)); -BER_ex = mean(bhat_ex(1:Lmin) ~= (b_te(1:Lmin)>0)); -fprintf('BER (ML-based MLSE): %.3e\n', BER_ml); -fprintf('BER (classic MLSE ): %.3e\n', BER_ex); -end - -% ----------------- helpers ----------------- -function Y = im2col_sliding(y_pad, Nwin, Delta) -T = length(y_pad) - (Nwin-1) - Delta; -Y = zeros(Nwin,T); -for k=1:T - Y(:,k) = y_pad(k-Delta : k-Delta+Nwin-1); -end -end - -function yw = feat_std(y_pad,k,Nwin,Delta,mu_win,sd_win) -yw = y_pad(k-Delta : k-Delta+Nwin-1); -yw = (yw - mu_win) ./ sd_win; -end - -function c = learnedBM(W,b0,yw,sp,s,~) -% pick parameters of the (sp->s) transition -persistent map; -if isempty(map) - % build once: index of W/b0 for each (sp,s) - nStates = size(W,1)*0+1; %#ok -end -% linear scan is fine at this size: -% (use first match of (sp,s)) -c = inf; -for t=1:size(W,2) - % suppose we stored (sp,s) order as in training; we cannot access here. - % Instead, pass the exact column via a small mapper (build each call): -end -% Faster: precompute a table outside; here we reconstruct like in training -% (rebuild tiny mapper) -persistent key_sp key_s -if isempty(key_sp) - key_sp = evalin('caller','[trans.sp]'); - key_s = evalin('caller','[trans.s]'); -end -t = find(key_sp==sp & key_s==s,1); -z = W(:,t).'*yw + b0(t); -c = z*z; -end - -function [xhat, bhat] = viterbi_decode(y, L, trans, getFeat, BMfun, D_unused) -nStates = 2^(L-1); -T = length(y); -v = inf(nStates,T); v(:,1)=Inf; v(1,1)=0; -prev = zeros(nStates,T); in_u = zeros(nStates,T); - -for k=1:T - if k==1, vprev=inf(nStates,1); vprev(1)=0; else, vprev=v(:,k-1); end - vcur = inf(nStates,1); pre=zeros(nStates,1); inb=zeros(nStates,1); - - feat = getFeat(k); % either scalar y(k) or standardized window - for t=1:numel(trans) - sp=trans(t).sp; s=trans(t).s; u=trans(t).u; - ck = BMfun(feat, sp, s, u); - cand = vprev(sp)+ck; - if cand < vcur(s) - vcur(s)=cand; pre(s)=sp; inb(s)=u; - end - end - v(:,k)=vcur; prev(:,k)=pre; in_u(:,k)=inb; -end - -[~,st]=min(v(:,T)); -xhat=zeros(T,1); -for k=T:-1:1 - xhat(k)=in_u(st,k); - st=prev(st,k); if st==0, st=1; end -end -bhat = xhat>0; -end diff --git a/projects/ML_based_MLSE/model.m b/projects/ML_based_MLSE/model.m index dd8c018..4610b72 100644 --- a/projects/ML_based_MLSE/model.m +++ b/projects/ML_based_MLSE/model.m @@ -18,13 +18,12 @@ laser_linewidth = 1e6; % Channel link_length = 0; -alpha = 0; doub_mode = db_mode.no_db; cols = linspecer(6); -rop = [-5]; +rop = [-8]; bwl = [0.5:0.1:1.5]; -fsym = [212:16:256].*1e9; +fsym = [160:16:256].*1e9; % nonlin_mod = [0.5:0.01:0.75]; nonlin_mod = ones(size(fsym)).*0.5; @@ -100,6 +99,7 @@ for r = 1:length(fsym) %% mu_lms = 0.0005; + pf_ncoeffs = 2; eq_ = FFE("epochs_tr",5,"epochs_dd",5,"len_tr",2^13,"mu_dd",mu_lms,"mu_tr",mu_lms,"order",50,"sps",2,"dd_mode",1,"adaption_technique","lms"); pf_ = Postfilter("ncoeff",pf_ncoeffs,"useBurg",1); @@ -116,49 +116,63 @@ for r = 1:length(fsym) % Sequence Est mlse_ = MLSE("duobinary_output",0,'M',M,'trellis_states',PAMmapper(M,0).levels,'scale_mode',0,'trellis_exclusion',0,'trellis_state_mode',2,'debug',0,'DIR',pf_.coefficients); [y_mlse] = mlse_.process(y_white,Symbols); - - Vit_bits = PAMmapper(M, 0, "eth_style", 0).demap(y_mlse); - [~, errors, ber_mlse_normal, errpos] = calc_ber(Vit_bits.signal, Tx_bits.signal, "skip_front", 10, "skip_end", 10, "returnErrorLocation", 1); + mlse_bits = PAMmapper(M, 0, "eth_style", 0).demap(y_mlse); + [~, errors, ber_mlse_normal, errpos] = calc_ber(mlse_bits.signal, Tx_bits.signal, "skip_front", 10, "skip_end", 10, "returnErrorLocation", 1); fprintf('MLSE: %.2e \n',ber_mlse_normal); showLevelHistogram(y_ffe,Symbols,"displayname",'ffe','fignum',111); - % showLevelHistogram(y_white,Symbols,"displayname",'ffe','fignum',111); - - %% optimize length - - mu_lms = 0.15; - eq = ML_MLSE("epochs_tr",5,"epochs_dd",5,"len_tr",2^14,... - "mu_dd",mu_lms,"mu_tr",mu_lms,"order",5,"sps",2,... - "traceback_depth",128,"L",2,"delta",0); - - [y_ml_mlse,Vit_signal] = eq.process(Rx_sig_2sps,Symbols); - y_ml_mlse_ = y_ml_mlse; - y_ml_mlse_.signal = circshift(y_ml_mlse.signal,0); - Vit_bits = PAMmapper(M, 0, "eth_style", 0).demap(y_ml_mlse_); - [~, errors, ber, errpos] = calc_ber(Vit_bits.signal, Tx_bits.signal, "skip_front", 10, "skip_end", 10, "returnErrorLocation", 1); - fprintf('ML MLSE: %.2e \n',ber); bursts = count_error_bursts(errpos, 10); - e = zeros(size(Vit_bits.signal)); + e = zeros(size(mlse_bits.signal)); e(errpos) = 1; figure(8) stem(e) + %% RUN ML-Based MLSE + + mu_lms = 0.15; + ml_mlse_equalizer = ML_MLSE("epochs_tr",50,"epochs_dd",10,"len_tr",Rx_sig_2sps.length-100,... + "mu_dd",mu_lms,"mu_tr",mu_lms,"order",15,"sps",2,... + "traceback_depth",128,"L",3,"delta",5); + + %% + ml_mlse_equalizer.epochs_tr = 50; + ml_mlse_equalizer.epochs_dd = 1; + [y_ml_mlse,Vit_signal] = ml_mlse_equalizer.process(Rx_sig_2sps,Symbols); + ml_mlse_bits = PAMmapper(M, 0, "eth_style", 0).demap(y_ml_mlse); + [~, errors, ber, errpos] = calc_ber(ml_mlse_bits.signal, Tx_bits.signal, "skip_front", 10, "skip_end", 10, "returnErrorLocation", 1); + fprintf('ML MLSE BER: %.2e \n',ber); + + bursts = count_error_bursts(errpos, 10); + e = zeros(size(ml_mlse_bits.signal)); + e(errpos) = 1; + figure(8) + stem(e) + + figure() + plot(ml_mlse_equalizer.ber) + beautifyBERplot + + + + + + %% optimize delta deltas = [-1:4]; ber = zeros(1,length(deltas)); parfor m = 1:numel(deltas) mu_lms = 0.2; - eq = ML_MLSE("epochs_tr",2,"epochs_dd",5,"len_tr",2^13,... + ml_mlse_equalizer = ML_MLSE("epochs_tr",2,"epochs_dd",5,"len_tr",2^13,... "mu_dd",mu_lms,"mu_tr",mu_lms,"order",4,"sps",1,... "traceback_depth",128,"L",3,"delta",deltas(m)); - [y_ml_mlse,Vit_signal] = eq.process(y_ffe,Symbols); + [y_ml_mlse,Vit_signal] = ml_mlse_equalizer.process(y_ffe,Symbols); y_ml_mlse_ = y_ml_mlse; y_ml_mlse_.signal = circshift(y_ml_mlse.signal,0); - Vit_bits = PAMmapper(M, 0, "eth_style", 0).demap(y_ml_mlse_); - [~, errors, ber(m), errpos] = calc_ber(Vit_bits.signal, Tx_bits.signal, "skip_front", 0, "skip_end", 0, "returnErrorLocation", 1); + mlse_bits = PAMmapper(M, 0, "eth_style", 0).demap(y_ml_mlse_); + [~, errors, ber(m), errpos] = calc_ber(mlse_bits.signal, Tx_bits.signal, "skip_front", 0, "skip_end", 0, "returnErrorLocation", 1); fprintf('ML MLSE: %.2e \n',ber(m)); end diff --git a/projects/ML_based_MLSE/rate_evaluation.m b/projects/ML_based_MLSE/rate_evaluation.m new file mode 100644 index 0000000..07e4057 --- /dev/null +++ b/projects/ML_based_MLSE/rate_evaluation.m @@ -0,0 +1,109 @@ + +ber_ffe = []; +ber_mlse = []; +ber_dbtgt = []; +ber_ml = []; + +mlse = 1; +dbtgt = 1; + +baudrates = [136:8:224].*1e9; +parfor i = 1:length(baudrates) + + rop = -8; + M = 4; + [Rx_sig_2sps_v1, Symbols_v1, Tx_bits_v1] = standard_link_model("M",M,"fsym",baudrates(i),"rop",rop,"laser_linewidth",1300,"link_length_m",0,"random_key",1,"apply_pulsef",0); + % [Rx_sig_2sps_v2, Symbols_v2, Tx_bits_v2] = standard_link_model("M",M,"fsym",200e9,"rop",rop,"laser_linewidth",1300,"link_length_m",0,"random_key",2); + % [Rx_sig_2sps_v3, Symbols_v3, Tx_bits_v3] = standard_link_model("M",M,"fsym",200e9,"rop",rop,"laser_linewidth",1300,"link_length_m",0,"random_key",3); + + %% FFE + MLSE + if mlse + pf_ncoeffs = 1; + ffe_order = [50, 0, 0]; + mu_ffe = [0.0001, 0.0008, 0.001]; + mu_dfe = 0.0004; + eq_ = EQ("Ne",ffe_order,"Nb",[0,0,0],"training_length",2^13,"training_loops",5,"dd_loops",5,"K",2,"DCmu",0.005,"DDmu",[mu_ffe mu_dfe],"DFEmu",0.005,"FFEmu",0,"plotfinal",0,"ideal_dfe",1); + pf_ = Postfilter("ncoeff",pf_ncoeffs,"useBurg",1); + mlse_ = MLSE("duobinary_output",0,'M',M,'trellis_states',PAMmapper(M,0).levels,'scale_mode',2,'trellis_exclusion',0,'trellis_state_mode',2,'debug',0,'DIR',pf_.coefficients); + [ffe_results, mlse_results] = vnle_postfilter_mlse(eq_, pf_, mlse_, M, Rx_sig_2sps_v1, Symbols_v1, Tx_bits_v1, ... + "precode_mode", duob_mode,... + 'showAnalysis', 0, ... + "postFFE", [],... + "eth_style_symbol_mapping", 0); + + ber_ffe(i) = ffe_results.metrics.BER; + ber_mlse(i) = mlse_results.metrics.BER; + end + + + %% FFE DB tgt. + MLSE + if dbtgt + mlse_db_ = MLSE("DIR",[1,1],"duobinary_output",0,"M",M,"trellis_states",PAMmapper(M,0).levels,'scale_mode',2,'trellis_exclusion',0,'trellis_state_mode',3); + ffe_order = [50, 0, 0]; + mu_ffe = [0.0001, 0.0008, 0.001]; + mu_dfe = 0.0004; + eq_ = EQ("Ne",ffe_order,"Nb",[0,0,0],"training_length",2^13,"training_loops",5,"dd_loops",5,"K",2,"DCmu",0.005,"DDmu",[mu_ffe mu_dfe],"DFEmu",0.005,"FFEmu",0,"plotfinal",0,"ideal_dfe",1); + + dbt_results = duobinary_target(eq_, mlse_db_, M, Rx_sig_2sps_v1, Symbols_v1, Tx_bits_v1, ... + "precode_mode", duob_mode, ... + 'showAnalysis', 0,... + "postFFE", []); + + ber_dbtgt(i) = dbt_results.metrics.BER; + end + + + %% + mu_lms = 0.0005; + pf_ncoeffs = 2; + eq_ = FFE("epochs_tr",5,"epochs_dd",5,"len_tr",2^13,"mu_dd",mu_lms,"mu_tr",mu_lms,"order",50,"sps",2,"dd_mode",1,"adaption_technique","lms"); + pf_ = Postfilter("ncoeff",pf_ncoeffs,"useBurg",1); + + % FFE + [y_ffe, ffe_noise] = eq_.process(Rx_sig_2sps_v1, Symbols_v1); + + + % Postfilter + [y_white,whitened_noise] = pf_.process(y_ffe, ffe_noise); + + %% RUN ML-Based MLSE + + mu_lms = 0.15; + ml_mlse_equalizer = ML_MLSE("epochs_tr",30,"epochs_dd",1,"len_tr",2^15,... + "mu_dd",mu_lms,"mu_tr",mu_lms,"order",5,"sps",1,... + "traceback_depth",128,"L",3,"delta",0); + + [y_ml_mlse,~] = ml_mlse_equalizer.process(y_white,Symbols_v1); + ml_mlse_bits = PAMmapper(M, 0, "eth_style", 0).demap(y_ml_mlse); + [~, errors, ber_ml(i), errpos] = calc_ber(ml_mlse_bits.signal, Tx_bits_v1.signal, "skip_front", 10, "skip_end", 10, "returnErrorLocation", 1); + fprintf('ML MLSE BER: %.2e \n',ber_ml(i)); + + % figure(11);hold on + % plot(1:numel(ml_mlse_equalizer.ber),ml_mlse_equalizer.ber); + % beautifyBERplot; + % xlim([1,numel(ml_mlse_equalizer.ber)]) + +end + +%% + +figure(6); hold on; +if mlse +plot(baudrates,ber_ffe,'DisplayName','FFE'); +plot(baudrates,ber_mlse,'DisplayName','MLSE'); +end +if dbtgt +plot(baudrates,ber_dbtgt,'DisplayName','DB tgt'); +end +plot(baudrates,ber_ml,'DisplayName','ML-MLSE'); +beautifyBERplot; +legend + + + + + + + + + diff --git a/projects/ML_based_MLSE/rop_evaluation.m b/projects/ML_based_MLSE/rop_evaluation.m new file mode 100644 index 0000000..68ed3eb --- /dev/null +++ b/projects/ML_based_MLSE/rop_evaluation.m @@ -0,0 +1,95 @@ + +ber_ffe = []; +ber_mlse = []; +ber_dbtgt = []; +ber_ml = []; + +mlse = 1; +dbtgt = 1; + +rops = linspace(-15,-5,12); +parfor i = 1:length(rops) + + rop = rops(i); + M = 4; + [Rx_sig_2sps_v1, Symbols_v1, Tx_bits_v1] = standard_link_model("M",M,"fsym",224e9,"rop",rop,"laser_linewidth",1300,"link_length_m",0,"random_key",1); + % [Rx_sig_2sps_v2, Symbols_v2, Tx_bits_v2] = standard_link_model("M",M,"fsym",200e9,"rop",rop,"laser_linewidth",1300,"link_length_m",0,"random_key",2); + % [Rx_sig_2sps_v3, Symbols_v3, Tx_bits_v3] = standard_link_model("M",M,"fsym",200e9,"rop",rop,"laser_linewidth",1300,"link_length_m",0,"random_key",3); + + %% FFE + MLSE + if mlse + pf_ncoeffs = 1; + ffe_order = [50, 0, 0]; + mu_ffe = [0.0001, 0.0008, 0.001]; + mu_dfe = 0.0004; + eq_ = EQ("Ne",ffe_order,"Nb",[0,0,0],"training_length",2^13,"training_loops",5,"dd_loops",5,"K",2,"DCmu",0.005,"DDmu",[mu_ffe mu_dfe],"DFEmu",0.005,"FFEmu",0,"plotfinal",0,"ideal_dfe",1); + pf_ = Postfilter("ncoeff",pf_ncoeffs,"useBurg",1); + mlse_ = MLSE("duobinary_output",0,'M',M,'trellis_states',PAMmapper(M,0).levels,'scale_mode',2,'trellis_exclusion',0,'trellis_state_mode',2,'debug',0,'DIR',pf_.coefficients); + [ffe_results, mlse_results] = vnle_postfilter_mlse(eq_, pf_, mlse_, M, Rx_sig_2sps_v1, Symbols_v1, Tx_bits_v1, ... + "precode_mode", duob_mode,... + 'showAnalysis', 0, ... + "postFFE", [],... + "eth_style_symbol_mapping", 0); + + ber_ffe(i) = ffe_results.metrics.BER; + ber_mlse(i) = mlse_results.metrics.BER; + end + + + %% FFE DB tgt. + MLSE + if dbtgt + mlse_db_ = MLSE("DIR",[1,1],"duobinary_output",0,"M",M,"trellis_states",PAMmapper(M,0).levels,'scale_mode',2,'trellis_exclusion',0,'trellis_state_mode',3); + ffe_order = [50, 0, 0]; + mu_ffe = [0.0001, 0.0008, 0.001]; + mu_dfe = 0.0004; + eq_ = EQ("Ne",ffe_order,"Nb",[0,0,0],"training_length",2^13,"training_loops",5,"dd_loops",5,"K",2,"DCmu",0.005,"DDmu",[mu_ffe mu_dfe],"DFEmu",0.005,"FFEmu",0,"plotfinal",0,"ideal_dfe",1); + + dbt_results = duobinary_target(eq_, mlse_db_, M, Rx_sig_2sps_v1, Symbols_v1, Tx_bits_v1, ... + "precode_mode", duob_mode, ... + 'showAnalysis', 0,... + "postFFE", []); + + ber_dbtgt(i) = dbt_results.metrics.BER; + end + + %% RUN ML-Based MLSE + + mu_lms = 0.15; + ml_mlse_equalizer = ML_MLSE("epochs_tr",30,"epochs_dd",1,"len_tr",2^14,... + "mu_dd",mu_lms,"mu_tr",mu_lms,"order",4,"sps",2,... + "traceback_depth",128,"L",2,"delta",0); + + [y_ml_mlse,~] = ml_mlse_equalizer.process(Rx_sig_2sps_v1,Symbols_v1); + ml_mlse_bits = PAMmapper(M, 0, "eth_style", 0).demap(y_ml_mlse); + [~, errors, ber_ml(i), errpos] = calc_ber(ml_mlse_bits.signal, Tx_bits_v1.signal, "skip_front", 10, "skip_end", 10, "returnErrorLocation", 1); + fprintf('ML MLSE BER: %.2e \n',ber_ml(i)); + + % figure(11);hold on + % plot(1:numel(ml_mlse_equalizer.ber),ml_mlse_equalizer.ber); + % beautifyBERplot; + % xlim([1,numel(ml_mlse_equalizer.ber)]) + +end + +%% + +figure(3); hold on; +if mlse +plot(rops,ber_ffe,'DisplayName','FFE'); +plot(rops,ber_mlse,'DisplayName','MLSE'); +end +if dbtgt +plot(rops,ber_dbtgt,'DisplayName','DB tgt'); +end +plot(rops,ber_ml,'DisplayName','ML-MLSE'); +beautifyBERplot; +legend + + + + + + + + + diff --git a/projects/ML_based_MLSE/standard_link_model.m b/projects/ML_based_MLSE/standard_link_model.m new file mode 100644 index 0000000..2412af1 --- /dev/null +++ b/projects/ML_based_MLSE/standard_link_model.m @@ -0,0 +1,95 @@ +function [Rx_sig_2sps,Symbols,Tx_bits] = standard_link_model(options) + + % STANDARD_LINK_MODEL Basic IM/DD link simulation + % Rx_sig_2sps = standard_link_model(...optional args...) + % + % All arguments are optional and default to standard parameters + % if not provided. + + arguments + + % --- Transmitter settings --- + options.M (1,1) double = 4 + options.apply_pulsef (1,1) logical = true + options.fdac (1,1) double = 256e9 + options.fadc (1,1) double = 256e9 + options.random_key (1,1) double = 2 + options.rcalpha (1,1) double = 0.05 + options.kover (1,1) double = 8 + options.vbias_rel (1,1) double = 0.5 + options.u_pi (1,1) double = 3.2 + options.laser_wavelength (1,1) double = 1310 + options.laser_linewidth (1,1) double = 1e6 + + % --- Channel parameters --- + options.link_length_m (1,1) double = 0 + options.rop (1,:) double = -5 + options.fsym (1,:) double = (212:16:256)*1e9 + options.doub_mode (1,1) db_mode = db_mode.no_db + + % --- Debug --- + options.debug (1,1) logical = false + + end + + % --- Pulse former --- + Pform = Pulseformer("fsym",options.fsym,"fdac",4*options.fsym, ... + "pulse","rc","pulselength",16,"alpha",options.rcalpha); + + % --- Transmitter source --- + [Digi_sig,Symbols,Tx_bits] = PAMsource( ... + "fsym",options.fsym,"M",options.M,"order",18,"useprbs",0, ... + "fs_out",options.fdac,"applyclipping",0,"clipfactor",1.5, ... + "applypulseform",options.apply_pulsef,"pulseformer",Pform, ... + "randkey",options.random_key,"db_precode",0,"db_encode",0, ... + "mrds_code",0,"mrds_blocklength",512, ... + "duobinary_mode",options.doub_mode).process(); + + % --- AWG driver --- + El_sig = M8199B("kover",options.kover).process(Digi_sig); + El_sig = El_sig.normalize("mode","oneone"); + + % --- E/O Modulation --- + vbias = -options.vbias_rel*options.u_pi; + Opt_sig = EML("mode",eml_mode.im_cosinus,"power",3, ... + "fsimu",El_sig.fs,"lambda",options.laser_wavelength, ... + "bias",vbias,"u_pi",options.u_pi,"linewidth",options.laser_linewidth, ... + "randomkey",options.random_key+1).process(El_sig); + + % --- Fiber --- + Opt_sig = Fiber("fsimu",Opt_sig.fs,"fiber_length",options.link_length_m, ... + "alpha",0.3,"D",0,"lambda0",1310,"gamma",0,"Dslope",0.07).process(Opt_sig); + + % --- Amplifier (ROP set) --- + Opt_sig = Amplifier("amp_mode","ideal_no_noise", ... + "gain_mode","output_power","amplification_db",options.rop).process(Opt_sig); + + % --- Photodiode --- + PD_sig = Photodiode("fsimu",options.fdac*options.kover,"dark_current",2e-8, ... + "responsivity",1,"temperature",20,"nep",1.8e-11, ... + "randomkey",options.random_key).process(Opt_sig); + + % --- Electrical LPF (receiver frontend) --- + rx_bwl = 70e9; + PD_sig = Filter('filtdegree',4,"f_cutoff",rx_bwl, ... + "fs",options.fdac*options.kover,"filterType",filtertypes.butterworth, ... + "active",true).process(PD_sig); + + % --- Scope low-pass and sampling --- + Lp_scpe = Filter('filtdegree',4,"f_cutoff",110e9,"fs",options.fadc, ... + "filterType",filtertypes.butterworth,"active",true); + + Scpe_sig = Scope("fsimu",options.fdac*options.kover,"fadc",options.fadc, ... + "delay",0,"fixed_delay",0,"filtertype",filtertypes.butterworth, ... + "samplingdelay",0,"rand_samplingdelay",0,"freq_offset",0, ... + "samp_jitter",0,"adcresolution",8,"quantbuffer",0.1, ... + 'block_dc',1,'lpf_active',1,'H_lpf',Lp_scpe).process(PD_sig); + + % --- Downsample to 2 sps --- + Scpe_sig_2sps = Scpe_sig.resample("fs_out",2*options.fsym); + [~,Scpe_cell,~,found_sync] = Scpe_sig_2sps.tsynch( ... + "reference",Symbols,"fs_ref",options.fsym,"debug_plots",0); + + Rx_sig_2sps = Scpe_cell{1}.normalize("mode","rms"); + +end