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 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 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