348 lines
15 KiB
Matlab
348 lines
15 KiB
Matlab
classdef FFE_A2Residual < FFE_plain
|
|
%FFE_A2Residual Plain FFE with A2 residual MPI suppression.
|
|
%
|
|
% This class reuses FFE_plain for process flow and BER-based mu
|
|
% optimization. Only equalize is specialized to add the paper-style A2
|
|
% residual correction: average y_raw - d_hat per PAM level, subtract the
|
|
% level-weighted residual estimate, then redo the DD decision.
|
|
|
|
properties
|
|
dc_level_avg_bufferlength_a2
|
|
dc_level_update_blocklength_a2
|
|
dc_smoothing_a2
|
|
dc_level_weights_a2
|
|
end
|
|
|
|
methods
|
|
function obj = FFE_A2Residual(options)
|
|
arguments(Input)
|
|
options.sps = 2
|
|
options.order = 15
|
|
|
|
options.len_tr = 4096
|
|
options.mu_tr = 0
|
|
options.epochs_tr = 5
|
|
|
|
options.adaption_technique adaption_method = adaption_method.lms
|
|
options.dd_mode = 1
|
|
options.mu_dd = 1e-5
|
|
options.epochs_dd = 5
|
|
options.dd_len_fraction = 1
|
|
|
|
options.dc_level_avg_bufferlength_a2 = 0
|
|
options.dc_level_update_blocklength_a2 = 0
|
|
options.dc_smoothing_a2 = 0
|
|
options.dc_level_weights_a2 = 0
|
|
|
|
options.decide = false
|
|
|
|
options.save_debug = 0
|
|
options.optmize_mus = 0
|
|
options.mu_optimization_len = 2^15
|
|
options.plot_mu_optimization = 0
|
|
options.mu_optimization_fignum = 3010
|
|
end
|
|
|
|
obj@FFE_plain( ...
|
|
"sps",options.sps, ...
|
|
"order",options.order, ...
|
|
"len_tr",options.len_tr, ...
|
|
"mu_tr",options.mu_tr, ...
|
|
"epochs_tr",options.epochs_tr, ...
|
|
"adaption_technique",options.adaption_technique, ...
|
|
"dd_mode",options.dd_mode, ...
|
|
"mu_dd",options.mu_dd, ...
|
|
"epochs_dd",options.epochs_dd, ...
|
|
"dd_len_fraction",options.dd_len_fraction, ...
|
|
"decide",options.decide, ...
|
|
"save_debug",options.save_debug, ...
|
|
"optmize_mus",options.optmize_mus, ...
|
|
"mu_optimization_len",options.mu_optimization_len, ...
|
|
"plot_mu_optimization",options.plot_mu_optimization, ...
|
|
"mu_optimization_fignum",options.mu_optimization_fignum);
|
|
|
|
obj.dc_level_avg_bufferlength_a2 = floor(options.dc_level_avg_bufferlength_a2);
|
|
obj.dc_level_update_blocklength_a2 = floor(options.dc_level_update_blocklength_a2);
|
|
obj.dc_smoothing_a2 = min(max(options.dc_smoothing_a2,0),1);
|
|
obj.dc_level_weights_a2 = options.dc_level_weights_a2;
|
|
|
|
assert(obj.dc_level_avg_bufferlength_a2 >= 0);
|
|
assert(obj.dc_level_update_blocklength_a2 >= 0);
|
|
end
|
|
|
|
function [y,d_hat] = equalize(obj,x,d,mu,epochs,N,training,showviz)
|
|
arguments
|
|
obj
|
|
x
|
|
d
|
|
mu
|
|
epochs
|
|
N
|
|
training
|
|
showviz
|
|
end
|
|
|
|
unused_showviz = showviz; %#ok<NASGU>
|
|
x = x(:);
|
|
d = d(:);
|
|
N = obj.validSampleLength(N,x,d);
|
|
n_symbols = N / obj.sps;
|
|
y = zeros(n_symbols,1);
|
|
d_hat = zeros(n_symbols,1);
|
|
|
|
if n_symbols == 0
|
|
return
|
|
end
|
|
|
|
if isempty(obj.constellation)
|
|
obj.constellation = unique(d);
|
|
end
|
|
decision_constellation = obj.constellation(:);
|
|
|
|
x = [zeros(floor(obj.order/2),1); x; zeros(obj.order,1)];
|
|
lambda = mu;
|
|
mask = ones(obj.order,1);
|
|
maincursor_pos = ceil(length(obj.e)/2);
|
|
adaption_code = obj.adaptionCode();
|
|
|
|
if mu == 0 || (~obj.dd_mode && ~training)
|
|
epochs = 1;
|
|
end
|
|
|
|
dc_level_enabled = obj.dc_level_avg_bufferlength_a2 > 1 && any(obj.dc_level_weights_a2(:) ~= 0);
|
|
if dc_level_enabled
|
|
[dc_level_weight_by_level,n_levels] = obj.expandedLevelWeights(decision_constellation);
|
|
dc_level_buffer_len = obj.dc_level_avg_bufferlength_a2;
|
|
dc_level_err_buffer = NaN(n_levels,dc_level_buffer_len);
|
|
dc_level_err_buffer_pos_by_level = zeros(n_levels,1);
|
|
dc_level_err_sum_by_level = zeros(n_levels,1);
|
|
dc_level_buffer_valid_count_by_level = zeros(n_levels,1);
|
|
dc_level_mpi_est_by_level = zeros(n_levels,1);
|
|
dc_level_valid_count_by_level = zeros(n_levels,1);
|
|
dc_level_update_blocklength = obj.dc_level_update_blocklength_a2;
|
|
if dc_level_update_blocklength <= 0
|
|
dc_level_update_blocklength = dc_level_buffer_len;
|
|
end
|
|
dc_level_update_blocklength = max(1,floor(dc_level_update_blocklength));
|
|
dc_level_window_future_fraction = obj.dc_smoothing_a2; %#ok<NASGU>
|
|
end
|
|
|
|
debug_enabled = obj.save_debug;
|
|
if debug_enabled
|
|
obj.initializeDebug(n_symbols,training);
|
|
end
|
|
|
|
for epoch = 1:epochs
|
|
symbol = 0;
|
|
for sample = 1:obj.sps:N
|
|
symbol = symbol + 1;
|
|
grad = zeros(obj.order,1);
|
|
update = zeros(obj.order,1);
|
|
weight = 0;
|
|
dc_level_mpi_est = 0;
|
|
dc_level_weight = 0;
|
|
dc_level_valid_count = 0;
|
|
|
|
U = x(obj.order+sample-1:-1:sample);
|
|
y(symbol,1) = (obj.e.*mask).' * U;
|
|
y_raw = y(symbol);
|
|
|
|
if training
|
|
d_hat(symbol,1) = d(symbol);
|
|
[~,symbol_idx] = min(abs(d_hat(symbol) - decision_constellation));
|
|
dc_level_decision_level = decision_constellation(symbol_idx);
|
|
else
|
|
[~,symbol_idx] = min(abs(y_raw - decision_constellation));
|
|
d_hat(symbol,1) = decision_constellation(symbol_idx);
|
|
dc_level_decision_level = decision_constellation(symbol_idx);
|
|
end
|
|
|
|
dc_level_symbol_idx = symbol_idx;
|
|
if dc_level_enabled
|
|
dc_level_mpi_est = dc_level_mpi_est_by_level(dc_level_symbol_idx);
|
|
dc_level_valid_count = dc_level_valid_count_by_level(dc_level_symbol_idx);
|
|
dc_level_weight = dc_level_weight_by_level(dc_level_symbol_idx) * ...
|
|
min(dc_level_valid_count / dc_level_buffer_len,1);
|
|
|
|
mpi_err = y_raw - d_hat(symbol);
|
|
y(symbol,1) = y_raw - dc_level_weight * dc_level_mpi_est;
|
|
|
|
if training
|
|
d_hat(symbol,1) = d(symbol);
|
|
else
|
|
[~,symbol_idx] = min(abs(y(symbol) - decision_constellation));
|
|
d_hat(symbol,1) = decision_constellation(symbol_idx);
|
|
dc_level_decision_level = decision_constellation(symbol_idx);
|
|
end
|
|
|
|
[dc_level_err_buffer,dc_level_err_buffer_pos_by_level, ...
|
|
dc_level_err_sum_by_level,dc_level_buffer_valid_count_by_level, ...
|
|
dc_level_mpi_est_by_level,dc_level_valid_count_by_level] = ...
|
|
obj.updateResidualBuffer( ...
|
|
dc_level_err_buffer,dc_level_err_buffer_pos_by_level, ...
|
|
dc_level_err_sum_by_level,dc_level_buffer_valid_count_by_level, ...
|
|
dc_level_mpi_est_by_level,dc_level_valid_count_by_level, ...
|
|
mpi_err,dc_level_symbol_idx,symbol,dc_level_update_blocklength);
|
|
end
|
|
|
|
err = d_hat(symbol) - y(symbol);
|
|
|
|
if mu ~= 0 && (training || obj.dd_mode)
|
|
switch adaption_code
|
|
case 1
|
|
weight = mu / ((U.'*U) + eps);
|
|
grad = err * U;
|
|
update = grad * weight;
|
|
obj.e = obj.e + update;
|
|
|
|
case 2
|
|
weight = mu;
|
|
grad = err * U;
|
|
update = grad * weight;
|
|
obj.e = obj.e + update;
|
|
|
|
case 3
|
|
denom = lambda + U.' * obj.P * U;
|
|
k = (obj.P * U) / denom;
|
|
update = k * err;
|
|
obj.e = obj.e + update;
|
|
obj.P = (1/lambda) * (obj.P - k * (U.' * obj.P));
|
|
end
|
|
end
|
|
|
|
if debug_enabled && epoch == 1
|
|
obj.debug_struct.error_first_epoch(1,symbol) = err * err';
|
|
end
|
|
|
|
if debug_enabled && epoch == epochs
|
|
error_power = err * err';
|
|
update_power = update.'*update ./ (sqrt((obj.e.'*obj.e) / obj.order) + eps);
|
|
obj.debug_struct.error(1,symbol) = error_power;
|
|
obj.debug_struct.main_cursor(1,symbol) = abs(obj.e(maincursor_pos));
|
|
obj.debug_struct.mu_nlms(1,symbol) = weight;
|
|
obj.debug_struct.update_gradient(1,symbol) = grad.'*grad;
|
|
obj.debug_struct.dc_level_mpi_est(1,symbol) = dc_level_mpi_est;
|
|
obj.debug_struct.dc_level_weight(1,symbol) = dc_level_weight;
|
|
obj.debug_struct.dc_level_valid_count(1,symbol) = dc_level_valid_count;
|
|
obj.debug_struct.dc_level_symbol_idx(1,symbol) = dc_level_symbol_idx;
|
|
obj.debug_struct.dc_level_decision_level(1,symbol) = dc_level_decision_level;
|
|
obj.debug_struct.dc_level_y_raw(1,symbol) = y_raw;
|
|
|
|
if training
|
|
obj.debug_struct.error_tr(1,symbol) = error_power;
|
|
obj.debug_struct.update_tr(1,symbol) = update_power;
|
|
else
|
|
obj.debug_struct.error_dd(1,symbol) = error_power;
|
|
obj.debug_struct.update(1,symbol) = update_power;
|
|
end
|
|
end
|
|
end
|
|
end
|
|
end
|
|
end
|
|
|
|
methods (Access = private)
|
|
function N = validSampleLength(obj,N,x,d)
|
|
N_available = min(numel(x),numel(d) * obj.sps);
|
|
N = min(N,N_available);
|
|
N = obj.sps * floor(N / obj.sps);
|
|
N = max(0,N);
|
|
end
|
|
|
|
function adaption_code = adaptionCode(obj)
|
|
if obj.adaption_technique == adaption_method.nlms
|
|
adaption_code = 1;
|
|
elseif obj.adaption_technique == adaption_method.lms
|
|
adaption_code = 2;
|
|
elseif obj.adaption_technique == adaption_method.rls
|
|
adaption_code = 3;
|
|
else
|
|
builtin("error","FFE_A2Residual:InvalidAdaptionTechnique", ...
|
|
"Unsupported FFE adaption technique.");
|
|
end
|
|
end
|
|
|
|
function [weights,n_levels] = expandedLevelWeights(obj,decision_constellation)
|
|
n_levels = numel(decision_constellation);
|
|
if isscalar(obj.dc_level_weights_a2)
|
|
weights = repmat(obj.dc_level_weights_a2,n_levels,1);
|
|
elseif numel(obj.dc_level_weights_a2) == n_levels
|
|
weights = obj.dc_level_weights_a2(:);
|
|
else
|
|
builtin("error","FFE_A2Residual:InvalidDCLevelWeights", ...
|
|
"dc_level_weights_a2 must be scalar or have one entry per constellation level.");
|
|
end
|
|
end
|
|
|
|
function initializeDebug(obj,n_symbols,training)
|
|
obj.debug_struct.error = NaN(1,n_symbols);
|
|
obj.debug_struct.error_first_epoch = NaN(1,n_symbols);
|
|
obj.debug_struct.main_cursor = NaN(1,n_symbols);
|
|
obj.debug_struct.mu_nlms = NaN(1,n_symbols);
|
|
obj.debug_struct.update_gradient = NaN(1,n_symbols);
|
|
obj.debug_struct.dc_level_mpi_est = NaN(1,n_symbols);
|
|
obj.debug_struct.dc_level_weight = NaN(1,n_symbols);
|
|
obj.debug_struct.dc_level_valid_count = NaN(1,n_symbols);
|
|
obj.debug_struct.dc_level_symbol_idx = NaN(1,n_symbols);
|
|
obj.debug_struct.dc_level_decision_level = NaN(1,n_symbols);
|
|
obj.debug_struct.dc_level_y_raw = NaN(1,n_symbols);
|
|
|
|
if training
|
|
obj.debug_struct.error_tr = NaN(1,n_symbols);
|
|
obj.debug_struct.update_tr = NaN(1,n_symbols);
|
|
else
|
|
obj.debug_struct.error_dd = NaN(1,n_symbols);
|
|
obj.debug_struct.update = NaN(1,n_symbols);
|
|
end
|
|
end
|
|
end
|
|
|
|
methods (Static, Access = private)
|
|
function [buffer,buffer_pos_by_level,err_sum_by_level,buffer_valid_count_by_level, ...
|
|
mpi_est_by_level,valid_count_by_level] = updateResidualBuffer( ...
|
|
buffer,buffer_pos_by_level,err_sum_by_level,buffer_valid_count_by_level, ...
|
|
mpi_est_by_level,valid_count_by_level,new_value,symbol_idx,symbol, ...
|
|
update_blocklength)
|
|
|
|
buffer_len = size(buffer,2);
|
|
buffer_pos = buffer_pos_by_level(symbol_idx) + 1;
|
|
if buffer_pos > buffer_len
|
|
buffer_pos = 1;
|
|
end
|
|
buffer_pos_by_level(symbol_idx) = buffer_pos;
|
|
|
|
old_value = buffer(symbol_idx,buffer_pos);
|
|
if isfinite(old_value)
|
|
err_sum_by_level(symbol_idx) = err_sum_by_level(symbol_idx) - old_value;
|
|
buffer_valid_count_by_level(symbol_idx) = buffer_valid_count_by_level(symbol_idx) - 1;
|
|
end
|
|
|
|
if isfinite(new_value)
|
|
buffer(symbol_idx,buffer_pos) = new_value;
|
|
err_sum_by_level(symbol_idx) = err_sum_by_level(symbol_idx) + new_value;
|
|
buffer_valid_count_by_level(symbol_idx) = buffer_valid_count_by_level(symbol_idx) + 1;
|
|
else
|
|
buffer(symbol_idx,buffer_pos) = NaN;
|
|
end
|
|
|
|
if mod(symbol,update_blocklength) == 0
|
|
if update_blocklength == 1
|
|
valid_count_by_level(symbol_idx) = buffer_valid_count_by_level(symbol_idx);
|
|
if valid_count_by_level(symbol_idx) == 0
|
|
mpi_est_by_level(symbol_idx) = 0;
|
|
else
|
|
mpi_est_by_level(symbol_idx) = ...
|
|
err_sum_by_level(symbol_idx) / valid_count_by_level(symbol_idx);
|
|
end
|
|
else
|
|
valid_count_by_level = buffer_valid_count_by_level;
|
|
has_valid = valid_count_by_level > 0;
|
|
mpi_est_by_level(:) = 0;
|
|
mpi_est_by_level(has_valid) = ...
|
|
err_sum_by_level(has_valid) ./ valid_count_by_level(has_valid);
|
|
end
|
|
end
|
|
end
|
|
end
|
|
end
|