diff --git a/Functions/EQ_recipes/mpi_recipe_dev.m b/Functions/EQ_recipes/mpi_recipe_dev.m index 1911f6e..e056ff9 100644 --- a/Functions/EQ_recipes/mpi_recipe_dev.m +++ b/Functions/EQ_recipes/mpi_recipe_dev.m @@ -18,11 +18,11 @@ end run_conventional_ffe = 1; run_ma = 0; run_hpf = 0; -run_a2_tracked_levels = 1; -run_a2_residual = 1; -run_a1 = 1; +run_a2_tracked_levels = 0; +run_a2_residual = 0; +run_a1 = 0; run_tracking_adaptive = 1; -plot_output_signals = 0; +plot_output_signals = 1; %% Shared fixed EQ settings eq_sps = 2; @@ -283,7 +283,7 @@ if run_tracking_adaptive "dc_tracking_persistence_gain", 0, ... "dc_tracking_buffer_len", block_update, ... "optmize_mus", false, ... - "optimize_dc_tracking_params", true, ... + "optimize_dc_tracking_params", false, ... "dc_tracking_optimization_len", 2^15, ... "dc_tracking_optimization_max_evals", 30, ... "plot_mu_optimization", options.debug_plots, ... @@ -296,12 +296,12 @@ if run_tracking_adaptive "dc_tracking", "adaptive", block_update); output.(char(storageName)) = ffe_results; - % fignum = 106; - % - % eq_noise = equalized_signal - Symbols; - % dn = sprintf("FFE DCT; SIR: %d dB",options.dataTable.sir); - % showEQNoisePSD(eq_noise, "fignum", fignum, "displayname", dn,"colormode","diverging"); - % ylim([-70 -30]); + fignum = 106; + + eq_noise = equalized_signal - Symbols; + dn = sprintf("FFE DCT; SIR: %d dB",options.dataTable.sir); + showEQNoisePSD(eq_noise, "fignum", fignum, "displayname", dn,"colormode","diverging"); + ylim([-70 -30]); % % dn = sprintf("FFE only; SIR: %d dB",options.dataTable.sir); % showEQNoisePSD(eq_noise_ffe, "fignum", fignum+1, "displayname", dn,"colormode","diverging"); diff --git a/projects/Diss/MPI_revisit/algorithms/PLOT_mpi_reduction_db_baudrate_comparison_pam4_1000m.m b/projects/Diss/MPI_revisit/algorithms/PLOT_mpi_reduction_db_baudrate_comparison_pam4_1000m.m new file mode 100644 index 0000000..f336f8f --- /dev/null +++ b/projects/Diss/MPI_revisit/algorithms/PLOT_mpi_reduction_db_baudrate_comparison_pam4_1000m.m @@ -0,0 +1,505 @@ +%% BER over SIR for PAM4 algorithm and baudrate comparison at 1000 m MPI delay +% 1) gather block_update = 1 data from the database +% 2) clean data and keep the 1000 m delay regime +% 3) plot one BER-over-SIR figure grouped by algorithm and baudrate + +clear; clc; + +%% 1) Gather data + +studyName = "block_update_sweep"; +selectedBlockUpdate = 1; +selectedPamLevel = 4; +selectedPathLength = 1000; + +selectedSymbolrates = [96e9 112e9]; + +algorithmSelection = table( ... + ["dc_tracking"; ... + "plain_ffe"; ... + "a2_tracked_levels"; ... + "a2_residual"; ... + "a1_moving_average"], ... + [true; true; false; false; false], ... + 'VariableNames', ["algorithm", "enabled"]); +selectedAlgorithms = algorithmSelection.algorithm(algorithmSelection.enabled); + +useBoundedLines = true; +usePolyfit = true; +polyfitOrderMax = 4; +boundaryPolyfitOrderMax = 4; +maxBerForPlot = 0.1; +scatterKeepFraction = 0.2; +boundMode = "fitStd"; % "fitStd", "directStd", "fitCI", or "directCI" +confidenceLevel = 0.95; + +db = DBHandler( ... + "dataBase", "labor", ... + "type", "mysql", ... + "server", "192.168.178.192", ... + "user", "silas", ... + "password", "silas"); +db.refresh(); + +fp = QueryFilter(); +fp.where('Runs', 'interference_path_length', 'EQUALS', selectedPathLength); +fp.where('Runs', 'fiber_length', 'EQUALS', 0); +fp.where('Runs', 'db_mode', 'EQUALS', '"no_db"'); +fp.where('Runs', 'v_bias', 'EQUALS', 2.65); +fp.where('Runs', 'pam_level', 'EQUALS', selectedPamLevel); +% fp.where('MpiReductionResults', 'study_name', 'EQUALS', char(studyName)); +fp.where('MpiReductionResults', 'block_update', 'EQUALS', selectedBlockUpdate); + +selectedFields = { ... + 'Runs.run_id', ... + 'Runs.sir', ... + 'Runs.pam_level', ... + 'Runs.symbolrate', ... + 'Runs.interference_path_length', ... + 'Runs.power_mpi_interference', ... + 'Runs.power_pd_in', ... + 'Runs.v_bias', ... + 'Runs.loop_id', ... + 'MpiReductionResults.occurrence_idx', ... + 'MpiReductionResults.storage_name', ... + 'MpiReductionResults.algorithm', ... + 'MpiReductionResults.algorithm_variant', ... + 'MpiReductionResults.eq_class', ... + 'MpiReductionResults.block_update', ... + 'MpiReductionResults.BER', ... + 'MpiReductionResults.BER_precoded', ... + 'MpiReductionResults.SNR', ... + 'MpiReductionResults.GMI', ... + 'MpiReductionResults.AIR'}; +selectedFields = selectedFields(:); + +[rawData, query] = db.queryDB(fp, selectedFields); +disp(query); +fprintf("Fetched %d MPI reduction rows.\n", height(rawData)); + +%% 2) Clean data + +data = rawData; +numericFields = ["run_id", "sir", "pam_level", "symbolrate", ... + "interference_path_length", "occurrence_idx", "block_update", ... + "power_mpi_interference", "BER", "BER_precoded", "SNR", "GMI", "AIR"]; +for fieldIdx = 1:numel(numericFields) + fieldName = numericFields(fieldIdx); + if ismember(fieldName, string(data.Properties.VariableNames)) + data.(char(fieldName)) = numericColumn(data.(char(fieldName))); + end +end + +stringFields = ["storage_name", "algorithm", "algorithm_variant", "eq_class"]; +for fieldIdx = 1:numel(stringFields) + fieldName = stringFields(fieldIdx); + if ismember(fieldName, string(data.Properties.VariableNames)) + data.(char(fieldName)) = stringColumn(data.(char(fieldName))); + end +end + +data(data.run_id == 3866, :) = []; +data(data.run_id == 3865, :) = []; +data(data.run_id == 3796, :) = []; +data(data.run_id == 3797, :) = []; +data(data.run_id == 3798, :) = []; +data(data.run_id == 4002, :) = []; +data(data.run_id == 4200, :) = []; +data(data.run_id == 4199, :) = []; +data(data.loop_id == 91, :) = []; +data(data.loop_id == 59, :) = []; + +data = data(isfinite(data.BER) & data.BER > 0 & data.BER < maxBerForPlot, :); +data.sir_exact = -7 - data.power_mpi_interference; + +data = data(data.pam_level == selectedPamLevel, :); +data = data(data.interference_path_length == selectedPathLength, :); +data = data(ismember(data.symbolrate, selectedSymbolrates), :); +data = data(ismember(data.algorithm, selectedAlgorithms), :); +selectedAlgorithms = selectedAlgorithms(ismember(selectedAlgorithms, unique(data.algorithm, "stable"))); + +if isempty(data) + warning("PLOT_mpi_reduction_db_baudrate_comparison_pam4_1000m:NoRows", ... + "No rows remain after BER/study/block/PAM/symbolrate/algorithm filtering."); + return +end + +data.clean_keep = true(height(data), 1); +groupId = findgroups(data.symbolrate, data.algorithm, data.sir); +for curGroup = unique(groupId(isfinite(groupId))).' + rowMask = groupId == curGroup; + berValues = data.BER(rowMask); + if nnz(rowMask) > 3 + data.clean_keep(rowMask) = ~isoutlier(berValues); + end +end +cleanData = data(data.clean_keep, :); + +groupVars = ["symbolrate", "algorithm", "sir"]; +summaryTable = groupsummary(cleanData, groupVars, ... + {"mean", "min", "max"}, "BER"); +sirExactTable = groupsummary(cleanData, groupVars, "median", "sir_exact"); +summaryTable = sortrows(summaryTable, groupVars); +sirExactTable = sortrows(sirExactTable, groupVars); +summaryTable.sir_exact = sirExactTable.median_sir_exact; +summaryTable = addBerIntervalBounds(summaryTable, cleanData, groupVars, confidenceLevel); + +fprintf("Cleaned to %d rows across %d baudrate/algorithm/SIR groups.\n", ... + height(cleanData), height(summaryTable)); +disp(groupcounts(cleanData, ["symbolrate", "algorithm"])); + +%% 3) Plot BER-over-SIR algorithm comparison with baudrate curves + +figure(); hold on; + +for algIdx = 1:numel(selectedAlgorithms) + algorithmName = selectedAlgorithms(algIdx); + algColor = algorithmColor(algorithmName); + displayName = algorithmDisplayName(algorithmName); + + for rateIdx = 1:numel(selectedSymbolrates) + symbolrate = selectedSymbolrates(rateIdx); + [marker, lineStyle] = symbolrateStyle(symbolrate); + + rawMask = cleanData.symbolrate == symbolrate & ... + cleanData.algorithm == algorithmName; + curveMask = summaryTable.symbolrate == symbolrate & ... + summaryTable.algorithm == algorithmName; + + if ~any(curveMask) + continue + end + + scatterRows = downsampleRows(cleanData(rawMask, :), scatterKeepFraction); + % scatter(scatterRows.sir_exact, scatterRows.BER, ... + % 14, ... + % "Marker", marker, ... + % "MarkerEdgeColor", algColor, ... + % "MarkerFaceColor", algColor, ... + % "MarkerEdgeAlpha", 0.25, ... + % "MarkerFaceAlpha", 0.25, ... + % "HandleVisibility", "off"); + + sirValues = summaryTable.sir_exact(curveMask).'; + meanBer = summaryTable.mean_BER(curveMask).'; + [boundCenterBer, boundLowerBer, boundUpperBer] = ... + selectBerBounds(summaryTable(curveMask, :), boundMode); + valid = isfinite(sirValues) & isfinite(meanBer) & meanBer > 0; + + if useBoundedLines && exist("boundedline", "file") && any(valid) + [xBand, centerBand, yBounds] = berIntervalBounds( ... + sirValues, boundCenterBer, boundLowerBer, boundUpperBer, ... + boundMode, boundaryPolyfitOrderMax); + + [hl, hp] = boundedline(xBand, centerBand, yBounds, ... + 'alpha', 'transparency', 0.08, ... + 'cmap', algColor, ... + 'nan', 'fill', ... + 'orientation', 'vert'); + set(hl, "LineStyle", "none", "Marker", "none", "HandleVisibility", "off"); + set(hp, "LineStyle", "-", "HandleVisibility", "off", "Marker", "none"); + end + + plot(sirValues(valid), meanBer(valid), ... + "LineStyle", lineStyle, ... + "Marker", marker, ... + "MarkerSize", 3.5, ... + "LineWidth", 1, ... + "Color", algColor, ... + "MarkerFaceColor", "w", ... + "MarkerEdgeColor", algColor, ... + "DisplayName", sprintf("%s, %.0f GBd", displayName, symbolrate * 1e-9)); + + if usePolyfit + fitMask = valid & meanBer > 0; + if nnz(fitMask) >= 2 + fitOrder = min(polyfitOrderMax, nnz(fitMask) - 1); + fitCoeff = polyfit(sirValues(fitMask), log10(meanBer(fitMask)), fitOrder); + xFit = linspace(min(sirValues(fitMask)), max(sirValues(fitMask)), 300); + yFit = 10 .^ polyval(fitCoeff, xFit); + + % plot(xFit, yFit, ... + % "LineStyle", "--", ... + % "LineWidth", 1.1, ... + % "Color", algColor, ... + % "HandleVisibility", "off"); + end + end + end +end + +yline(2.2e-4, "LineWidth", 1, "LineStyle", "--", "HandleVisibility", "off"); +yline(3.8e-3, "LineWidth", 1, "LineStyle", "--", "HandleVisibility", "off"); +yline(2e-2, "LineWidth", 1, "LineStyle", "--", "HandleVisibility", "off"); + +% title(sprintf("PAM %.0f, %.0f m, block update %.0f", ... +% selectedPamLevel, selectedPathLength, selectedBlockUpdate)); +xlabel("SIR (dB)"); +ylabel("BER"); +set(gca, "YScale", "log"); +ylim([3e-5, maxBerForPlot]); +xlim([14, 45]); +grid on; +box on; +legend("Location", "northeast", "Interpreter", "none"); + +if exist("beautifyBERplot", "file") + beautifyBERplot("logscale", true, "setcolors", false, ... + "setmarkers", false, "changemarkers", false); +end + +% mat2tikz_improved("C:/Users/Silas/Documents/6971e0b65b380ca6d71c837f/04_Experimental_Evaluation/tikz/mpi/ber_vs_sir_pam4_baudrates_1000m.tikz","cleanfigure",1); + +%% Local helpers + +function values = numericColumn(values) + if iscell(values) + values = string(values); + end + if isstring(values) || ischar(values) + values = str2double(values); + end + values = double(values); +end + +function values = stringColumn(values) + if iscell(values) + values = string(values); + elseif ischar(values) + values = string(values); + end + values = strip(string(values)); +end + +function rows = downsampleRows(rows, keepFraction) + if keepFraction >= 1 || height(rows) <= 1 + return + end + + keepEvery = max(1, round(1 / keepFraction)); + keepIdx = 1:keepEvery:height(rows); + rows = rows(keepIdx, :); +end + +function summaryTable = addBerIntervalBounds(summaryTable, cleanData, groupVars, confidenceLevel) + nGroups = height(summaryTable); + ciCenter = NaN(nGroups, 1); + ciLower = NaN(nGroups, 1); + ciUpper = NaN(nGroups, 1); + stdCenter = NaN(nGroups, 1); + stdLower = NaN(nGroups, 1); + stdUpper = NaN(nGroups, 1); + + for groupIdx = 1:nGroups + rowMask = true(height(cleanData), 1); + for varIdx = 1:numel(groupVars) + varName = groupVars(varIdx); + rowMask = rowMask & cleanData.(char(varName)) == summaryTable.(char(varName))(groupIdx); + end + + [ciCenter(groupIdx), ciLower(groupIdx), ciUpper(groupIdx)] = ... + logBerMeanConfidenceInterval(cleanData.BER(rowMask), confidenceLevel); + [stdCenter(groupIdx), stdLower(groupIdx), stdUpper(groupIdx)] = ... + logBerMeanStdInterval(cleanData.BER(rowMask)); + end + + summaryTable.ci_center_BER = ciCenter; + summaryTable.ci_lower_BER = ciLower; + summaryTable.ci_upper_BER = ciUpper; + summaryTable.std_center_BER = stdCenter; + summaryTable.std_lower_BER = stdLower; + summaryTable.std_upper_BER = stdUpper; +end + +function [centerBer, lowerBer, upperBer] = logBerMeanConfidenceInterval(berValues, confidenceLevel) + berValues = berValues(isfinite(berValues) & berValues > 0); + if isempty(berValues) + centerBer = NaN; + lowerBer = NaN; + upperBer = NaN; + return + end + + logBer = log10(berValues(:)); + centerLog = mean(logBer, "omitnan"); + + if numel(logBer) < 2 + centerBer = 10 .^ centerLog; + lowerBer = centerBer; + upperBer = centerBer; + return + end + + alpha = 1 - confidenceLevel; + if exist("tinv", "file") + tCritical = tinv(1 - alpha/2, numel(logBer) - 1); + else + tCritical = 1.96; + end + + halfWidth = tCritical * std(logBer, 0, "omitnan") / sqrt(numel(logBer)); + centerBer = 10 .^ centerLog; + lowerBer = 10 .^ (centerLog - halfWidth); + upperBer = 10 .^ (centerLog + halfWidth); +end + +function [centerBer, lowerBer, upperBer] = logBerMeanStdInterval(berValues) + berValues = berValues(isfinite(berValues) & berValues > 0); + if isempty(berValues) + centerBer = NaN; + lowerBer = NaN; + upperBer = NaN; + return + end + + logBer = log10(berValues(:)); + centerLog = mean(logBer, "omitnan"); + stdLog = std(logBer, 0, "omitnan"); + + centerBer = 10 .^ centerLog; + lowerBer = 10 .^ (centerLog - stdLog); + upperBer = 10 .^ (centerLog + stdLog); +end + +function [centerBer, lowerBer, upperBer] = selectBerBounds(curveSummary, boundMode) + if boundMode == "fitCI" || boundMode == "directCI" + centerBer = curveSummary.ci_center_BER.'; + lowerBer = curveSummary.ci_lower_BER.'; + upperBer = curveSummary.ci_upper_BER.'; + else + centerBer = curveSummary.std_center_BER.'; + lowerBer = curveSummary.std_lower_BER.'; + upperBer = curveSummary.std_upper_BER.'; + end +end + +function [xBand, centerBand, yBounds] = berIntervalBounds(sirValues, centerBer, lowerBer, upperBer, boundMode, maxOrder) + if boundMode == "directCI" || boundMode == "directStd" + [xBand, centerBand, yBounds] = directBerBounds(sirValues, centerBer, lowerBer, upperBer); + return + end + + [xBand, centerBand, yBounds] = fittedBerBounds(sirValues, centerBer, lowerBer, upperBer, maxOrder); +end + +function [xBand, centerBand, yBounds] = directBerBounds(sirValues, centerBer, lowerBer, upperBer) + valid = isfinite(sirValues) & isfinite(centerBer) & isfinite(lowerBer) & ... + isfinite(upperBer) & centerBer > 0 & lowerBer > 0 & upperBer > 0; + + xBand = sirValues(valid).'; + centerBand = centerBer(valid).'; + lowerBand = lowerBer(valid).'; + upperBand = upperBer(valid).'; + + [xBand, orderIdx] = sort(xBand(:)); + centerBand = centerBand(orderIdx); + lowerBand = lowerBand(orderIdx); + upperBand = upperBand(orderIdx); + + lowerTmp = min(lowerBand, upperBand); + upperBand = max(lowerBand, upperBand); + lowerBand = lowerTmp; + + centerBand = min(max(centerBand, lowerBand), upperBand); + yBounds = [max(centerBand - lowerBand, 0), max(upperBand - centerBand, 0)]; +end + +function [xBand, centerBand, yBounds] = fittedBerBounds(sirValues, meanBer, minBer, maxBer, maxOrder) + valid = isfinite(sirValues) & isfinite(meanBer) & isfinite(minBer) & ... + isfinite(maxBer) & meanBer > 0 & minBer > 0 & maxBer > 0; + + x = sirValues(valid); + yMean = meanBer(valid); + yMin = minBer(valid); + yMax = maxBer(valid); + + if numel(x) < 2 + xBand = x(:); + centerBand = yMean(:); + yBounds = [max(yMean(:) - yMin(:), 0), max(yMax(:) - yMean(:), 0)]; + return + end + + [x, orderIdx] = sort(x(:)); + yMean = yMean(orderIdx); + yMin = yMin(orderIdx); + yMax = yMax(orderIdx); + + xBand = linspace(min(x), max(x), 300).'; + fitOrder = min(maxOrder, numel(unique(x)) - 1); + + if fitOrder < 1 + centerBand = interp1(x, yMean, xBand, "linear", "extrap"); + lowerBand = interp1(x, yMin, xBand, "linear", "extrap"); + upperBand = interp1(x, yMax, xBand, "linear", "extrap"); + else + centerBand = fitLogBer(x, yMean, xBand, fitOrder); + lowerBand = fitLogBer(x, yMin, xBand, fitOrder); + upperBand = fitLogBer(x, yMax, xBand, fitOrder); + end + + lowerTmp = min(lowerBand, upperBand); + upperBand = max(lowerBand, upperBand); + lowerBand = lowerTmp; + + centerBand = min(max(centerBand, lowerBand), upperBand); + yBounds = [max(centerBand - lowerBand, 0), max(upperBand - centerBand, 0)]; +end + +function yFit = fitLogBer(x, y, xFit, fitOrder) + coeff = polyfit(x, log10(y), fitOrder); + yFit = 10 .^ polyval(coeff, xFit); +end + +function label = algorithmDisplayName(algorithmName) + algorithmName = string(algorithmName); + switch algorithmName + case "plain_ffe" + label = "FFE only"; + case "a2_tracked_levels" + label = "ACT"; + case "a2_residual" + label = "L-DCA"; + case "a1_moving_average" + label = "DCA"; + case "dc_tracking" + label = "DCT"; + otherwise + label = algorithmName; + end +end + +function color = algorithmColor(algorithmName) + algorithmName = string(algorithmName); + switch algorithmName + case "plain_ffe" + color = [0.3467 0.5360 0.6907]; + case "a2_tracked_levels" + color = [0.9153 0.2816 0.2878]; + case "a2_residual" + color = [0.4416 0.7490 0.4322]; + case "a1_moving_average" + color = [1.0000 0.5984 0.2000]; + case "dc_tracking" + color = [0.6769 0.4447 0.7114]; + otherwise + color = [0 0 0]; + end +end + +function [marker, lineStyle] = symbolrateStyle(symbolrate) + switch symbolrate + case 72e9 + marker = "square"; + lineStyle = "--"; + case 96e9 + marker = "o"; + lineStyle = "-"; + case 112e9 + marker = "diamond"; + lineStyle = ":"; + otherwise + marker = "^"; + lineStyle = "-."; + end +end diff --git a/projects/Diss/MPI_revisit/investigate_mpi_algorithms.m b/projects/Diss/MPI_revisit/investigate_mpi_algorithms.m index 78ea07f..225bbdf 100644 --- a/projects/Diss/MPI_revisit/investigate_mpi_algorithms.m +++ b/projects/Diss/MPI_revisit/investigate_mpi_algorithms.m @@ -3,8 +3,8 @@ dsp_options = struct(); dsp_options.mode = "run_id"; dsp_options.recipe = @mpi_recipe_dev; dsp_options.append_to_db = false; -dsp_options.start_occurence = 1; -dsp_options.max_occurences = 15; +dsp_options.start_occurence = 5; +dsp_options.max_occurences = 1; dsp_options.debug_plots = false; write_mpi_reduction_db = 0; @@ -37,7 +37,6 @@ db = DBHandler("dataBase", [dsp_options.dataBase],... "user", dsp_options.user, "password", dsp_options.password); %% - pamformats = [4,6,8]; baudrates = [112e9,96e9,72e9]; @@ -48,10 +47,10 @@ for i = 1 fp = QueryFilter(); fp.where('Runs','run_id','GREATER_EQUAL',3153); - fp.where('Runs', 'symbolrate', 'EQUALS', B); % 72 96 112 + fp.where('Runs', 'symbolrate','EQUALS',B); % 72 96 112 fp.where('Runs', 'fiber_length', 'EQUALS', 0); fp.where('Runs', 'interference_path_length', 'EQUALS', 1000); - fp.where('Runs', 'sir', 'LESS_EQUAL', 50); + fp.where('Runs', 'sir', 'LESS_EQUAL', 20); fp.where('Runs', 'db_mode', 'EQUALS', '"no_db"'); % fp.where('Runs', 'is_mpi', 'EQUALS', 1); fp.where('Runs', 'pam_level', 'EQUALS', M); @@ -62,11 +61,12 @@ for i = 1 [~, sortIdx] = sort(dataTable.sir, 'descend'); dataTable = dataTable(sortIdx, :); - % desired_sir = [15,22,26,30,38]; - % desired_sir = [22,30,38]; + desired_sir = [15,22,26,30,38]; + desired_sir = [22,38]; + % exclude_sr = [112e9]; % % Filter dataTable to only keep rows with sir in desired_sir - % isDesired = ismember(dataTable.sir, desired_sir); - % dataTable = dataTable(isDesired, :); + isDesired = ~ismember(dataTable.symbolrate, desired_sir); + dataTable = dataTable(isDesired, :); % dataTable = dataTable(1,:); run_ids = dataTable.run_id;