Final MPI Calculations here

This commit is contained in:
Silas Oettinghaus
2026-07-14 11:50:22 +02:00
parent f421348e5b
commit 8c4edf2490
8 changed files with 1003 additions and 12 deletions

View File

@@ -11,14 +11,22 @@ classdef FFE_DCTracking < FFE_plain
dc_tracking_persistence_gain
dc_tracking_power_exponent
dc_tracking_buffer_len
dc_tracking_mu_eff_max
e_dc
optimize_dc_tracking_params = 0
dc_tracking_optimization
dc_tracking_optimization_iter = 0
dc_tracking_optimization_len
dc_tracking_optimization_max_evals
dc_tracking_optimization_delay_weight
dc_tracking_optimization_smoothing_len
end
properties (Access = private)
dc_tracking_mu_min = 1e-6
dc_tracking_mu_max = 3e-1
dc_tracking_mu_eff_min = 0
dc_tracking_mu_eff_max = inf
end
methods
@@ -42,6 +50,7 @@ classdef FFE_DCTracking < FFE_plain
options.dc_tracking_persistence_gain = 0
options.dc_tracking_power_exponent = 2
options.dc_tracking_buffer_len = 1
options.dc_tracking_mu_eff_max = inf
options.decide = false
@@ -50,6 +59,11 @@ classdef FFE_DCTracking < FFE_plain
options.mu_optimization_len = 2^15
options.plot_mu_optimization = 0
options.mu_optimization_fignum = 3010
options.optimize_dc_tracking_params = 0
options.dc_tracking_optimization_len = 2^15
options.dc_tracking_optimization_max_evals = 20
options.dc_tracking_optimization_delay_weight = 1e-3
options.dc_tracking_optimization_smoothing_len = 501
end
obj@FFE_plain( ...
@@ -75,8 +89,18 @@ classdef FFE_DCTracking < FFE_plain
obj.dc_tracking_persistence_gain = options.dc_tracking_persistence_gain;
obj.dc_tracking_power_exponent = options.dc_tracking_power_exponent;
obj.dc_tracking_buffer_len = floor(options.dc_tracking_buffer_len);
obj.dc_tracking_mu_eff_max = options.dc_tracking_mu_eff_max;
obj.e_dc = 0;
obj.optimize_dc_tracking_params = options.optimize_dc_tracking_params;
obj.dc_tracking_optimization_len = options.dc_tracking_optimization_len;
obj.dc_tracking_optimization_max_evals = max(1, ...
floor(options.dc_tracking_optimization_max_evals));
obj.dc_tracking_optimization_delay_weight = ...
max(0,options.dc_tracking_optimization_delay_weight);
obj.dc_tracking_optimization_smoothing_len = max(1, ...
floor(options.dc_tracking_optimization_smoothing_len));
assert(obj.dc_tracking_buffer_len >= 0);
end
@@ -90,6 +114,11 @@ classdef FFE_DCTracking < FFE_plain
obj.resetTrackingState();
end
if obj.optimize_dc_tracking_params
obj.optimizeDcTrackingParams(X.signal,D.signal);
obj.resetTrackingState();
end
training = true;
showviz = false;
obj.equalize(X.signal,D.signal,obj.mu_tr,obj.epochs_tr,obj.len_tr,training,showviz);
@@ -134,7 +163,8 @@ classdef FFE_DCTracking < FFE_plain
if obj.dd_mode
vars = [vars, optimizableVariable("mu_dd",mu_range,"Transform","log")];
end
optimize_dc_tracking_mu = obj.dc_tracking_mu ~= 0;
optimize_dc_tracking_mu = obj.dc_tracking_mu ~= 0 && ...
~obj.optimize_dc_tracking_params;
if optimize_dc_tracking_mu
vars = [vars, optimizableVariable("dc_tracking_mu",[1e-5,1e-1],"Transform","log")];
end
@@ -175,6 +205,61 @@ classdef FFE_DCTracking < FFE_plain
end
end
function optimizeDcTrackingParams(obj,x,d)
[x_opt,d_opt,N_opt] = obj.optimizationSignals( ...
x,d,obj.dc_tracking_optimization_len);
vars = optimizableVariable( ...
"dc_tracking_mu",[1e-5,1e-1],"Transform","log");
initial_values = min(max(obj.dc_tracking_mu,1e-5),1e-1);
initial_names = "dc_tracking_mu";
if obj.dc_tracking_adaptive_enabled
vars = [vars, ...
optimizableVariable("dc_tracking_persistence_gain",[0,2]), ...
optimizableVariable("dc_tracking_mu_eff_max",[1e-3,3e-1], ...
"Transform","log")];
initial_values = [initial_values, ...
min(max(obj.dc_tracking_persistence_gain,0),2), ...
min(max(obj.dc_tracking_mu_eff_max,1e-3),3e-1)];
initial_names = [initial_names, ...
"dc_tracking_persistence_gain", "dc_tracking_mu_eff_max"];
end
initial_x = array2table(initial_values, ...
"VariableNames",cellstr(initial_names));
obj.dc_tracking_optimization_iter = 0;
fprintf("FFE_DCTracking DC opt uses fixed mu_tr=%9.3e, mu_dd=%9.3e on %d samples / %d symbols\n", ...
obj.mu_tr,obj.mu_dd,N_opt,numel(d_opt));
old_rng = rng;
cleanup_rng = onCleanup(@()rng(old_rng));
rng(42,"twister");
obj.dc_tracking_optimization = bayesopt( ...
@(p)obj.dcTrackingObjective(p,x_opt,d_opt),vars, ...
"MaxObjectiveEvaluations",obj.dc_tracking_optimization_max_evals, ...
"InitialX",initial_x, ...
"AcquisitionFunctionName","expected-improvement-plus", ...
"IsObjectiveDeterministic",true, ...
"Verbose",0, ...
"PlotFcn",[]);
clear cleanup_rng
best = obj.dc_tracking_optimization.XAtMinObjective;
obj.dc_tracking_mu = best.dc_tracking_mu;
if obj.dc_tracking_adaptive_enabled
obj.dc_tracking_persistence_gain = ...
best.dc_tracking_persistence_gain;
obj.dc_tracking_mu_eff_max = best.dc_tracking_mu_eff_max;
fprintf("\nFFE_DCTracking DC opt done: dc_tracking_mu=%9.3e, persistence_gain=%6.3f, mu_eff_max=%9.3e, objective=%9.3e\n", ...
obj.dc_tracking_mu,obj.dc_tracking_persistence_gain, ...
obj.dc_tracking_mu_eff_max, ...
obj.dc_tracking_optimization.MinObjective);
else
fprintf("\nFFE_DCTracking DC opt done: dc_tracking_mu=%9.3e, objective=%9.3e\n", ...
obj.dc_tracking_mu,obj.dc_tracking_optimization.MinObjective);
end
end
function objective = muObjective(obj,params,x,d)
old_state = obj.captureTrackingObjectiveState();
cleanup = onCleanup(@()obj.restoreTrackingObjectiveState(old_state));
@@ -218,6 +303,49 @@ classdef FFE_DCTracking < FFE_plain
clear cleanup
end
function objective = dcTrackingObjective(obj,params,x,d)
old_state = obj.captureTrackingObjectiveState();
cleanup = onCleanup(@()obj.restoreTrackingObjectiveState(old_state));
obj.applyDcTrackingObjectiveParams(params);
obj.save_debug = 1;
obj.resetTrackingState();
N_tr = min(obj.len_tr,numel(x));
obj.equalize(x,d,obj.mu_tr,obj.epochs_tr,N_tr,true,false);
if obj.dd_mode
[signal,~] = obj.equalize( ...
x,d,obj.mu_dd,obj.epochs_dd,numel(x),false,false);
else
[signal,~] = obj.equalize(x,d,0,1,numel(x),false,false);
end
[ber,errors] = obj.berObjective(signal,d);
[delay_symbols,delay_corr] = obj.dcTrackingDelayObjective(signal,d);
delay_penalty = obj.dc_tracking_optimization_delay_weight * ...
abs(delay_symbols) / max(numel(d),1);
objective = ber + delay_penalty;
if ~isfinite(objective)
objective = inf;
end
obj.dc_tracking_optimization_iter = ...
obj.dc_tracking_optimization_iter + 1;
if obj.dc_tracking_adaptive_enabled
fprintf("\rFFE_DCTracking DC opt %02d: dc_tracking_mu=%9.3e, persistence_gain=%6.3f, mu_eff_max=%9.3e, BER=%9.3e, delay=%7.0f, corr=%6.3f, obj=%9.3e, errors=%d", ...
obj.dc_tracking_optimization_iter,params.dc_tracking_mu, ...
params.dc_tracking_persistence_gain, ...
params.dc_tracking_mu_eff_max,ber,delay_symbols,delay_corr, ...
objective,errors);
else
fprintf("\rFFE_DCTracking DC opt %02d: dc_tracking_mu=%9.3e, BER=%9.3e, delay=%7.0f, corr=%6.3f, obj=%9.3e, errors=%d", ...
obj.dc_tracking_optimization_iter,params.dc_tracking_mu, ...
ber,delay_symbols,delay_corr,objective,errors);
end
clear cleanup
end
function [y,d_hat] = equalize(obj,x,d,mu,epochs,N,training,showviz)
arguments
obj
@@ -392,6 +520,9 @@ classdef FFE_DCTracking < FFE_plain
state.save_debug = obj.save_debug;
state.debug_struct = obj.debug_struct;
state.dc_tracking_mu = obj.dc_tracking_mu;
state.dc_tracking_persistence_gain = ...
obj.dc_tracking_persistence_gain;
state.dc_tracking_mu_eff_max = obj.dc_tracking_mu_eff_max;
end
function restoreTrackingObjectiveState(obj,state)
@@ -403,6 +534,106 @@ classdef FFE_DCTracking < FFE_plain
obj.save_debug = state.save_debug;
obj.debug_struct = state.debug_struct;
obj.dc_tracking_mu = state.dc_tracking_mu;
obj.dc_tracking_persistence_gain = ...
state.dc_tracking_persistence_gain;
obj.dc_tracking_mu_eff_max = state.dc_tracking_mu_eff_max;
end
function applyDcTrackingObjectiveParams(obj,params)
var_names = string(params.Properties.VariableNames);
if any(var_names == "dc_tracking_mu")
obj.dc_tracking_mu = params.dc_tracking_mu;
end
if any(var_names == "dc_tracking_persistence_gain")
obj.dc_tracking_persistence_gain = ...
params.dc_tracking_persistence_gain;
end
if any(var_names == "dc_tracking_mu_eff_max")
obj.dc_tracking_mu_eff_max = params.dc_tracking_mu_eff_max;
end
end
function [delay_symbols,delay_corr] = ...
dcTrackingDelayObjective(obj,signal,d)
delay_symbols = 0;
delay_corr = 0;
if ~isfield(obj.debug_struct,"dc_tracking_est") || ...
isempty(obj.debug_struct.dc_tracking_est)
return
end
smooth_len = obj.dc_tracking_optimization_smoothing_len;
dc_tracking_est_s = movmean( ...
obj.debug_struct.dc_tracking_est(:),smooth_len,"omitnan");
avg_lvl_dc = obj.averageLevelTrace(signal,d,smooth_len);
xcorr_len = min(numel(dc_tracking_est_s),numel(avg_lvl_dc));
if xcorr_len < 2
return
end
inv_dc_xcorr = -dc_tracking_est_s(1:xcorr_len);
avg_lvl_xcorr = avg_lvl_dc(1:xcorr_len);
inv_dc_xcorr = fillmissing( ...
inv_dc_xcorr,"linear","EndValues","nearest");
avg_lvl_xcorr = fillmissing( ...
avg_lvl_xcorr,"linear","EndValues","nearest");
inv_dc_xcorr = inv_dc_xcorr - mean(inv_dc_xcorr,"omitnan");
avg_lvl_xcorr = avg_lvl_xcorr - mean(avg_lvl_xcorr,"omitnan");
if rms(inv_dc_xcorr) <= eps || rms(avg_lvl_xcorr) <= eps
return
end
[dc_level_xcorr,dc_level_lags] = xcorr( ...
inv_dc_xcorr,avg_lvl_xcorr,"coeff");
[delay_corr,delay_idx] = max(dc_level_xcorr);
delay_symbols = dc_level_lags(delay_idx);
if ~isfinite(delay_corr)
delay_corr = 0;
delay_symbols = 0;
end
end
function avg_lvl_dc = averageLevelTrace(obj,signal,d,smooth_len)
signal = signal(:);
d = d(:);
n_symbols = min(numel(signal),numel(d));
signal = signal(1:n_symbols);
d = d(1:n_symbols);
levels = unique(d);
avg_for_lvl = NaN(numel(levels),n_symbols);
for level_idx = 1:numel(levels)
level_mask = d == levels(level_idx);
level_samples = signal(level_mask);
if isempty(level_samples)
continue
end
smooth_window = min(smooth_len,numel(level_samples));
avg_for_lvl(level_idx,level_mask) = movmean( ...
level_samples,smooth_window,"omitnan","Endpoints","shrink");
avg_for_lvl(level_idx,:) = obj.interpolateMissingAverage( ...
avg_for_lvl(level_idx,:));
end
avg_lvl_dc = mean(avg_for_lvl,1,"omitnan").';
end
function level_average = interpolateMissingAverage(~,level_average)
valid_samples = isfinite(level_average);
if nnz(valid_samples) == 0
return
elseif nnz(valid_samples) == 1
level_average(:) = level_average(valid_samples);
return
end
t = 1:numel(level_average);
level_average(~valid_samples) = interp1( ...
t(valid_samples),level_average(valid_samples), ...
t(~valid_samples),"linear","extrap");
end
function N = validSampleLength(obj,N,x,d)