Final MPI Calculations here
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user