第【51期】--基于IEEE 802.11a的OFDM深度学习均衡器设计--matlab完整代码
关注我,追更更多通信仿真!
文章目录
摘要
随着无线通信系统向更高频段、更复杂传播环境演进,传统信道均衡算法在性能与复杂度之间面临越来越严峻的权衡。本文基于IEEE 802.11a标准的OFDM物理层框架,构建了一套完整的端到端仿真平台,系统对比了基于卷积神经网络(CNN)的智能均衡器与传统迫零(ZF)和最小均方误差(MMSE)线性均衡器在不同信道条件下的误码率性能。仿真覆盖了AWGN、Rician(莱斯)衰落等多种场景,重点考察了莱斯K因子对均衡性能的影响。实验结果表明,CNN均衡器在Rician衰落信道下可获得相比MMSE更优的BER增益,且在高K因子条件下优势更为显著。本研究为深度学习技术在物理层接收机中的工程化应用提供了可复现的参考基准。
1 背景意义
1.1 从经典均衡到智能接收机
正交频分复用(OFDM)技术因其出色的抗多径能力和频谱效率,成为4G/5G及未来无线通信系统的核心物理层方案。然而,OFDM系统的性能在很大程度上取决于接收端信道均衡的质量。传统均衡器主要分为两类:迫零(ZF)均衡器通过完全消除信道间干扰来实现检测,但在深衰落子载波上会严重放大噪声;最小均方误差(MMSE)均衡器在干扰消除与噪声抑制之间寻求最优折中,是目前工程应用中最广泛的选择。
近年来,深度学习技术为物理层通信带来了范式层面的变革。研究表明,深度神经网络能够从大量数据中学习信道的统计特征和非线性失真模式,在无需精确信道先验信息的情况下实现高性能的信号检测与均衡。其中,卷积神经网络(CNN)凭借其对时频网格结构的天然适配性,在OFDM信号处理中展现出独特的优势。
1.2 意义
本文基于IEEE 802.11a标准构建了一套完整的OFDM仿真链路,从帧生成、前导码同步、信道估计到均衡检测,完整复现了真实接收机的信号处理流程。在此框架下,对CNN均衡器与ZF、MMSE进行了全面的性能对比,系统分析了SNR、信道场景、莱斯K因子等因素对均衡性能的影响,为深度学习均衡器的工程化评估提供了可复现的参考基准。
2 理论基础
2.1 OFDM系统与IEEE 802.11a帧结构
2.1.1 OFDM基本原理

2.1.2 IEEE 802.11a 帧结构
本文严格遵循 802.11a 规范,关键参数如下:
标准的802.11a物理层帧由前导码和数据两部分组成。前导码包含短训练序列(L-STF)用于粗同步和AGC设定,以及长训练序列(L-LTF)用于精细同步和信道估计。数据部分由多个OFDM符号构成,每个符号包含经过QPSK调制的有效载荷和固定模式的导频信号。导频信号不仅用于辅助信道估计,还可用于跟踪残留的载波相位误差。完整 PPDU 帧包含:
-
短训练序列(L-STF):10 个重复短符号(总长 8 μs),用于粗同步、AGC 和频偏估计。
-
长训练序列(L-LTF):两个重复长符号(总长 8 μs),用于精细信道估计。
-
数据字段:本文设置 50 个 OFDM 符号,承载 QPSK 未编码有效载荷(关闭 FEC 以纯粹考察均衡器性能)。
2.2 Rician 衰落模型


2.3 线性均衡器

2.3.1 迫零(ZF)均衡器

2.3.2 最小均方误差(MMSE)均衡器

在工程实践中,MMSE 均衡器因其对噪声的鲁棒性而成为大多数 OFDM 系统的默认选择,也是本文作为与CNN 均衡器进行性能对比的基线方法。
2.4 CNN均衡器
传统线性均衡器依赖精确的信道估计和显式的线性逆滤波,在信道估计误差、射频非线性和相位噪声等实际损伤下性能会显著下降。卷积神经网络(CNN)均衡器采用数据驱动的端到端学习范式,直接从接收的时频网格中学习发送符号的非线性映射关系,无需精确的信道先验。CNN在OFDM信号处理中具有天然优势:卷积核适配子载波的局部频域相关性,权重共享机制参数高效且具备平移等变性。
2.4.1 输入特征

2.4.2 网络结构

2.4.3 训练和推理过程

- 训练采用分阶段课程学习策略,从高信噪比、高 K 因子的简单样本逐步过渡到低信噪比、低 K 因子的困难样本,使网络能够稳健地建立从易到难的均衡能力。
- 推理阶段,取 Softmax 输出中最大概率对应的类别作为判决结果,再通过预定义的 Gray 映射表将类别索引转换为比特序列,与真实发送比特逐位比较以统计误码率
3 仿真流程设计
3.1 系统框架

整个仿真链路严格按照 IEEE 802.11a 物理层标准构建,涵盖发射机、信道、接收机同步与均衡检测三大模块。整体架构如上图所示。
3.2 仿真参数
| OFDM参数名称 | 符号/取值 | 说明 |
|---|---|---|
| FFT 点数 | N = 64 N = 64 N=64 | IEEE 802.11a 标准 |
| 循环前缀长度 | N g = 16 N_g = 16 Ng=16 采样点 | 对应 0.8 μs |
| 子载波间隔 | Δ f = 312.5 \Delta f = 312.5 Δf=312.5 kHz | = 20 MHz / 64 = 20\text{ MHz} / 64 =20 MHz/64 |
| 有用符号周期 | T u = 3.2 T_u = 3.2 Tu=3.2 μs | = 1 / Δ f = 1/\Delta f =1/Δf |
| 完整符号周期 | T s y m = 4.0 T_{sym} = 4.0 Tsym=4.0 μs | T u + T g T_u + T_g Tu+Tg |
| 有效子载波总数 | N a c t i v e = 52 N_{active} = 52 Nactive=52 | 48 数据 + 4 导频 |
| 数据子载波数 | N d a t a = 48 N_{data} = 48 Ndata=48 | 承载 QPSK 有效载荷 |
| 导频子载波数 | N p i l o t = 4 N_{pilot} = 4 Npilot=4 | 位置: − 21 , − 7 , 7 , 21 -21, -7, 7, 21 −21,−7,7,21 |
| 每帧 OFDM 符号数 | S = 50 S = 50 S=50 | QPSK 未编码数据 |
| 调制阶数 | M = 4 M = 4 M=4 | QPSK,每符号 2 bit |
| 前导码结构 | L-STF + L-LTF | 短/长训练序列 |
| 信道编码 | 关闭(FEC off) | 纯粹评估均衡器性能 |
| CNN参数 | 参数名称 | 取值 | 说明 |
|---|---|---|---|
| 输入层 | 输入通道数 | C = 10 C = 10 C=10 | 特征构造见 2.4.1 节 |
| 输入维度 | C × K C \times K C×K | K = 52 K=52 K=52 有效子载波 | |
| 卷积层 | 卷积核大小 | 9 | 一维卷积 |
| 扩张率序列 | [ 1 , 2 , 4 , 8 , 1 ] [1, 2, 4, 8, 1] [1,2,4,8,1] | 5 个残差块对应 | |
| 隐藏通道数 | 64 | 残差块内部特征维 | |
| 输出层 | 输出类别数 | M = 4 M = 4 M=4 | QPSK 四类符号 |
| 输出激活函数 | Softmax | 概率归一化 | |
| 激活函数 | 隐藏层激活 | ReLU | 除输出层外统一使用 |
| 残差结构 | 残差块数量 | 5 | 含跳跃连接 |
| 优化器 | 类型 | Adam | 自适应矩估计 |
| 初始学习率 | 8 × 10 − 4 8 \times 10^{-4} 8×10−4 | 课程学习首阶段 | |
| 最终学习率 | 3 × 10 − 5 3 \times 10^{-5} 3×10−5 | 课程学习末阶段 | |
| 训练配置 | 批次大小 | 256 | 每批训练样本数 |
| 每轮样本数 | 1.65 × 10 5 1.65 \times 10^5 1.65×105 | 各阶段固定 | |
| 损失函数 | 交叉熵 | 仅数据子载波 |
3.3 仿真图分析


可以看到:
- 在固定两径信道下扫描K因子,绝大部分情况下CNN大部分下均优于MMSE ,体现了CNN对不同信道条件的鲁棒性。
针对工程落地,单独统计了CNN处理单个帧(50个OFDM符号)的耗时(运行于Intel Xeon CPU,无GPU加速):

可以看到:
- cdf曲线越陡峭,说明延迟越集中,利于实时处理。
- CPU 测试结果显示端到端平均约 2.55ms,大于 0.2ms, 802.11a 处理需求;表明当前纯 CPU 实现尚无法满足实时处理要求,但若采用 GPU 加速、模型量化或减少每帧符号数,则推理延迟可大幅降低,具备实时部署潜力。
部分代码:
function report = run_cnn_deep_test(varargin)
% 多条件压力测试:在共享帧上对比 CNN MMSE 的性能。
parser = inputParser();
addParameter(parser, 'ModelPath', "");
addParameter(parser, 'EbNoDb', [0 5 10 15 20 25 30 35]);
addParameter(parser, 'Scenarios', ["AWGN" "LOS" "NLOS"]);
addParameter(parser, 'LOSKFactors', [3 8 15 30]);
addParameter(parser, 'NLOSKFactors', [1e-4 0.1 1]);
addParameter(parser, 'DopplerHz', [0 10 40]);
addParameter(parser, 'FramesPerCondition', 20);
addParameter(parser, 'RandomSeed', 271828);
addParameter(parser, 'WarmupIterations', 3);
addParameter(parser, 'SynchronizationThreshold', 0.05);
addParameter(parser, 'ConfidenceLevel', 0.95);
addParameter(parser, 'MakePlots', true);
addParameter(parser, 'Verbose', true);
parse(parser, varargin{:});
cfg = parser.Results;
cfg.Scenarios = upper(string(cfg.Scenarios(:).'));
validate_configuration(cfg);
rng(cfg.RandomSeed, 'twister');
baseParams = simulation_parameters().refreshDerived();
carrierMap = baseParams.getCarrierMap();
frameTool = IEEE80211aFrame(baseParams);
[cnn, modelPath, checkpointSelection] = ...
load_cnn_checkpoint(baseParams, cfg.ModelPath);
conditions = build_conditions(cfg);
numConditions = height(conditions);
methodNames = ["CNN" "ZF" "MMSE"];
numMethods = numel(methodNames);
totalFrameSamples = numConditions * cfg.FramesPerCondition;
bitErrors = zeros(numConditions, numMethods);
symbolErrors = zeros(numConditions, numMethods);
frameErrors = zeros(numConditions, numMethods);
totalBits = zeros(numConditions, 1);
totalSymbols = zeros(numConditions, 1);
syncFailures = zeros(numConditions, 1);
pilotTrackingFrames = zeros(numConditions, 1);
cnnFeatureSeconds = zeros(totalFrameSamples, 1);
cnnInferenceSeconds = zeros(totalFrameSamples, 1);
cnnPostprocessSeconds = zeros(totalFrameSamples, 1);
cnnEndToEndSeconds = zeros(totalFrameSamples, 1);
zfDetectorSeconds = zeros(totalFrameSamples, 1);
mmseDetectorSeconds = zeros(totalFrameSamples, 1);
timingConditionIndex = zeros(totalFrameSamples, 1);
warm_up_detector(cnn, baseParams, carrierMap, frameTool, ...
conditions(1, :), cfg);
bitsPerFrame = baseParams.NumPayloadCarriers * ...
baseParams.NumOFDMSymbols * log2(baseParams.ModulationOrder);
fprintf("=== CNN 多条件压力测试 ===\n");
fprintf("模型: %s\n", modelPath);
fprintf("条件数: %d | 每条件帧数: %d | 总帧数: %d\n", ...
numConditions, cfg.FramesPerCondition, totalFrameSamples);
fprintf("预期评估比特数: %.3g\n", ...
totalFrameSamples * bitsPerFrame);
timingIdx = 0;
for conditionIdx = 1:numConditions
condition = conditions(conditionIdx, :);
params = condition_parameters(baseParams, condition);
params.assertOFDMChannelConsistency(condition_label(condition));
for frameIdx = 1:cfg.FramesPerCondition
rng(frame_seed(cfg.RandomSeed, conditionIdx, frameIdx), 'twister');
[txWaveform, txFrame] = frameTool.createRandomFrame();
rxWaveform = multipath_channel(txWaveform, params);
rxFrame = frameTool.receive(rxWaveform);
truthClasses = cnn.symbolsToClasses(txFrame.PayloadSymbols);
truthBits = txFrame.PayloadBits;
[cnnClasses, cnnBits, cnnTiming] = detect_with_cnn( ...
cnn, rxFrame, txFrame, carrierMap, ...
baseParams.NumOFDMSymbols);
[zfClasses, zfBits, zfSeconds] = detect_with_linear_equalizer( ...
frameTool, rxFrame, 'ZF', baseParams.ModulationOrder);
[mmseClasses, mmseBits, mmseSeconds] = ...
detect_with_linear_equalizer(frameTool, rxFrame, ...
'MMSE', baseParams.ModulationOrder);
frameBitErrors = [ ...
sum(cnnBits(:) ~= truthBits(:)), ...
sum(zfBits(:) ~= truthBits(:)), ...
sum(mmseBits(:) ~= truthBits(:))];
frameSymbolErrors = [ ...
sum(cnnClasses(:) ~= truthClasses(:)), ...
sum(zfClasses(:) ~= truthClasses(:)), ...
sum(mmseClasses(:) ~= truthClasses(:))];
bitErrors(conditionIdx, :) = ...
bitErrors(conditionIdx, :) + frameBitErrors;
symbolErrors(conditionIdx, :) = ...
symbolErrors(conditionIdx, :) + frameSymbolErrors;
frameErrors(conditionIdx, :) = ...
frameErrors(conditionIdx, :) + (frameBitErrors > 0);
totalBits(conditionIdx) = totalBits(conditionIdx) + ...
numel(truthBits);
totalSymbols(conditionIdx) = totalSymbols(conditionIdx) + ...
numel(truthClasses);
syncFailures(conditionIdx) = syncFailures(conditionIdx) + ...
(rxFrame.SynchronizationMetric < ...
cfg.SynchronizationThreshold);
pilotTrackingFrames(conditionIdx) = ...
pilotTrackingFrames(conditionIdx) + ...
double(rxFrame.PilotPhaseTrackingEnabled);
timingIdx = timingIdx + 1;
cnnFeatureSeconds(timingIdx) = cnnTiming.FeatureSeconds;
cnnInferenceSeconds(timingIdx) = cnnTiming.InferenceSeconds;
cnnPostprocessSeconds(timingIdx) = cnnTiming.PostprocessSeconds;
cnnEndToEndSeconds(timingIdx) = cnnTiming.EndToEndSeconds;
zfDetectorSeconds(timingIdx) = zfSeconds;
mmseDetectorSeconds(timingIdx) = mmseSeconds;
timingConditionIndex(timingIdx) = conditionIdx;
end
if cfg.Verbose
conditionBER = bitErrors(conditionIdx, :) ./ ...
totalBits(conditionIdx);
fprintf("%3d/%3d | %-38s | BER CNN %.3e | ZF %.3e | MMSE %.3e\n", ...
conditionIdx, numConditions, condition_label(condition), ...
conditionBER(1), conditionBER(2), conditionBER(3));
end
end
zValue = normal_quantile(0.5 + cfg.ConfidenceLevel / 2);
ber = bitErrors ./ totalBits;
ser = symbolErrors ./ totalSymbols;
fer = frameErrors ./ cfg.FramesPerCondition;
[berLower, berUpper] = wilson_interval(bitErrors, totalBits, zValue);
[ferLower, ferUpper] = wilson_interval( ...
frameErrors, cfg.FramesPerCondition, zValue);
conditionTable = build_condition_results( ...
conditions, methodNames, bitErrors, symbolErrors, frameErrors, ...
totalBits, ber, ser, fer, berLower, berUpper, ...
ferLower, ferUpper, syncFailures, pilotTrackingFrames, ...
cfg.FramesPerCondition, timingConditionIndex, ...
cnnInferenceSeconds, cnnEndToEndSeconds);
summaryTable = aggregate_summary( ...
methodNames, bitErrors, symbolErrors, frameErrors, ...
totalBits, totalSymbols, cfg.FramesPerCondition, zValue);
scenarioSummary = grouped_summary( ...
conditions, methodNames, bitErrors, symbolErrors, frameErrors, ...
totalBits, totalSymbols, cfg.FramesPerCondition, zValue, ...
'Scenario');
snrSummary = grouped_summary( ...
conditions, methodNames, bitErrors, symbolErrors, frameErrors, ...
totalBits, totalSymbols, cfg.FramesPerCondition, zValue, ...
'ScenarioEbNo');
timingSummary = build_timing_summary( ...
cnnFeatureSeconds, cnnInferenceSeconds, cnnPostprocessSeconds, ...
cnnEndToEndSeconds, zfDetectorSeconds, mmseDetectorSeconds, ...
baseParams.NumOFDMSymbols);
comparison = build_comparison( ...
conditionTable, summaryTable, methodNames);
report = struct();
report.ModelPath = modelPath;
report.CheckpointSelection = checkpointSelection;
report.Configuration = cfg;
report.Parameters = baseParams;
report.MethodNames = methodNames;
report.Conditions = conditionTable;
report.Summary = summaryTable;
report.ScenarioSummary = scenarioSummary;
report.SNRSummary = snrSummary;
report.Comparison = comparison;
report.Timing = struct();
report.Timing.Summary = timingSummary;
report.Timing.CNNFeatureSeconds = cnnFeatureSeconds;
report.Timing.CNNInferenceSeconds = cnnInferenceSeconds;
report.Timing.CNNPostprocessSeconds = cnnPostprocessSeconds;
report.Timing.CNNEndToEndSeconds = cnnEndToEndSeconds;
report.Timing.ZFDetectorSeconds = zfDetectorSeconds;
report.Timing.MMSEDetectorSeconds = mmseDetectorSeconds;
report.Timing.ConditionIndex = timingConditionIndex;
report.Persistence = "已禁用: 结果仅保存在内存中";
print_summary(report);
if cfg.MakePlots
report.FigureHandles = plot_report(report);
end
end
function conditions = build_conditions(cfg)
% 根据配置生成所有测试条件(场景、K因子、多普勒、Eb/N0)
scenarioColumn = strings(0, 1);
kFactorColumn = zeros(0, 1);
dopplerColumn = zeros(0, 1);
ebNoColumn = zeros(0, 1);
for scenario = cfg.Scenarios
switch scenario
case "AWGN"
kFactors = NaN;
dopplerValues = 0;
case "LOS"
kFactors = double(cfg.LOSKFactors(:).');
dopplerValues = double(cfg.DopplerHz(:).');
case "NLOS"
kFactors = double(cfg.NLOSKFactors(:).');
dopplerValues = double(cfg.DopplerHz(:).');
end
for kFactor = kFactors
for dopplerHz = dopplerValues
for ebNoDb = double(cfg.EbNoDb(:).')
scenarioColumn(end+1, 1) = scenario; %#ok<AGROW>
kFactorColumn(end+1, 1) = kFactor; %#ok<AGROW>
dopplerColumn(end+1, 1) = dopplerHz; %#ok<AGROW>
ebNoColumn(end+1, 1) = ebNoDb; %#ok<AGROW>
end
end
end
end
conditions = table(scenarioColumn, kFactorColumn, dopplerColumn, ...
ebNoColumn, 'VariableNames', ...
{'Scenario', 'KFactor', 'DopplerHz', 'EbNoDb'});
end
function params = condition_parameters(baseParams, condition)
% 根据条件生成信道参数
if condition.Scenario == "AWGN"
params = configure_test_channel( ...
baseParams, condition.EbNoDb, condition.Scenario, 15);
else
params = configure_test_channel( ...
baseParams, condition.EbNoDb, condition.Scenario, 15, ...
'KFactor', condition.KFactor, ...
'MaximumDopplerShift', condition.DopplerHz);
end
end
function warm_up_detector(cnn, baseParams, carrierMap, frameTool, ...
condition, cfg)
% 预热CNN,确保首次推理不包含JIT编译开销
if cfg.WarmupIterations == 0
return;
end
params = condition_parameters(baseParams, condition);
rng(frame_seed(cfg.RandomSeed, 0, 1), 'twister');
[txWaveform, txFrame] = frameTool.createRandomFrame();
rxFrame = frameTool.receive(multipath_channel(txWaveform, params));
for warmupIdx = 1:cfg.WarmupIterations
detect_with_cnn(cnn, rxFrame, txFrame, carrierMap, ...
baseParams.NumOFDMSymbols);
end
end
function [classes, bits, timing] = detect_with_cnn( ...
cnn, rxFrame, txFrame, carrierMap, batchCount)
% 使用CNN进行检测,并计时各阶段
featureClock = tic;
X = build_frame_cnn_batch( ...
rxFrame, txFrame, carrierMap, cnn.NumInputChannels);
X = cnn.normalizeInputBatch(X, carrierMap.ActiveGlobalIdx);
featureSeconds = toc(featureClock);
inferenceClock = tic;
scores = gather(extractdata(forward( ...
cnn.Net, dlarray(single(X), 'CTB'))));
inferenceSeconds = toc(inferenceClock);
postprocessClock = tic;
scores = normalize_score_shape(scores, cnn, batchCount);
classes = cnn.oneHotToClasses(scores);
classes = classes(carrierMap.PayloadGlobalIdx, :);
bits = cnn.classesToBits(classes);
postprocessSeconds = toc(postprocessClock);
timing = struct();
timing.FeatureSeconds = featureSeconds;
timing.InferenceSeconds = inferenceSeconds;
timing.PostprocessSeconds = postprocessSeconds;
timing.EndToEndSeconds = ...
featureSeconds + inferenceSeconds + postprocessSeconds;
end
function [classes, bits, elapsedSeconds] = detect_with_linear_equalizer( ...
frameTool, rxFrame, method, modOrder)
% 使用线性均衡器(ZF或MMSE)进行检测并计时
detectorClock = tic;
equalized = frameTool.equalize(rxFrame, method);
qamIndices = qamdemod(equalized.PayloadSymbols, modOrder, 'gray', ...
'UnitAveragePower', true, 'OutputType', 'integer');
classes = double(qamIndices) + 1;
bitsColumn = qamdemod(equalized.PayloadSymbols(:), modOrder, 'gray', ...
'UnitAveragePower', true, 'OutputType', 'bit');
bits = reshape(uint8(bitsColumn), ...
[log2(modOrder), size(equalized.PayloadSymbols)]);
elapsedSeconds = toc(detectorClock);
end
function conditionTable = build_condition_results( ...
conditions, methodNames, bitErrors, symbolErrors, frameErrors, ...
totalBits, ber, ser, fer, berLower, berUpper, ...
ferLower, ferUpper, syncFailures, pilotTrackingFrames, ...
framesPerCondition, timingConditionIndex, ...
cnnInferenceSeconds, cnnEndToEndSeconds)
% 构建条件结果表
conditionTable = conditions;
conditionTable.Frames = repmat(framesPerCondition, height(conditions), 1);
conditionTable.TotalBits = totalBits;
conditionTable.SyncFailures = syncFailures;
conditionTable.PilotTrackingRate = ...
pilotTrackingFrames / framesPerCondition;
for methodIdx = 1:numel(methodNames)
prefix = char(methodNames(methodIdx));
conditionTable.([prefix 'BitErrors']) = bitErrors(:, methodIdx);
conditionTable.([prefix 'BER']) = ber(:, methodIdx);
conditionTable.([prefix 'BERLower']) = berLower(:, methodIdx);
conditionTable.([prefix 'BERUpper']) = berUpper(:, methodIdx);
conditionTable.([prefix 'SymbolErrors']) = ...
symbolErrors(:, methodIdx);
conditionTable.([prefix 'SER']) = ser(:, methodIdx);
conditionTable.([prefix 'FrameErrors']) = frameErrors(:, methodIdx);
conditionTable.([prefix 'FER']) = fer(:, methodIdx);
conditionTable.([prefix 'FERLower']) = ferLower(:, methodIdx);
conditionTable.([prefix 'FERUpper']) = ferUpper(:, methodIdx);
end
winners = strings(height(conditions), 1);
for conditionIdx = 1:height(conditions)
bestBER = min(ber(conditionIdx, :));
winners(conditionIdx) = strjoin( ...
methodNames(ber(conditionIdx, :) == bestBER), "=");
end
conditionTable.BestBERMethod = winners;
inferenceMeanMs = zeros(height(conditions), 1);
inferenceP95Ms = zeros(height(conditions), 1);
endToEndMeanMs = zeros(height(conditions), 1);
endToEndP95Ms = zeros(height(conditions), 1);
for conditionIdx = 1:height(conditions)
mask = timingConditionIndex == conditionIdx;
inferenceValues = 1000 * cnnInferenceSeconds(mask);
endToEndValues = 1000 * cnnEndToEndSeconds(mask);
inferenceMeanMs(conditionIdx) = mean(inferenceValues);
inferenceP95Ms(conditionIdx) = percentile(inferenceValues, 95);
endToEndMeanMs(conditionIdx) = mean(endToEndValues);
endToEndP95Ms(conditionIdx) = percentile(endToEndValues, 95);
end
conditionTable.CNNInferenceMeanMs = inferenceMeanMs;
conditionTable.CNNInferenceP95Ms = inferenceP95Ms;
conditionTable.CNNEndToEndMeanMs = endToEndMeanMs;
conditionTable.CNNEndToEndP95Ms = endToEndP95Ms;
end
function summaryTable = aggregate_summary( ...
methodNames, bitErrors, symbolErrors, frameErrors, totalBits, ...
totalSymbols, framesPerCondition, zValue)
% 汇总所有条件下的总体性能
totalFrameCount = size(bitErrors, 1) * framesPerCondition;
aggregateBitErrors = sum(bitErrors, 1).';
aggregateSymbolErrors = sum(symbolErrors, 1).';
aggregateFrameErrors = sum(frameErrors, 1).';
aggregateBits = sum(totalBits);
aggregateSymbols = sum(totalSymbols);
[berLower, berUpper] = wilson_interval( ...
aggregateBitErrors, aggregateBits, zValue);
[ferLower, ferUpper] = wilson_interval( ...
aggregateFrameErrors, totalFrameCount, zValue);
summaryTable = table(methodNames(:), aggregateBitErrors, ...
repmat(aggregateBits, numel(methodNames), 1), ...
aggregateBitErrors / aggregateBits, berLower, berUpper, ...
aggregateSymbolErrors, ...
aggregateSymbolErrors / aggregateSymbols, ...
aggregateFrameErrors, ...
aggregateFrameErrors / totalFrameCount, ferLower, ferUpper, ...
'VariableNames', {'Method', 'BitErrors', 'TotalBits', 'BER', ...
'BERLower', 'BERUpper', 'SymbolErrors', 'SER', 'FrameErrors', ...
'FER', 'FERLower', 'FERUpper'});
end
function result = grouped_summary( ...
conditions, methodNames, bitErrors, symbolErrors, frameErrors, ...
totalBits, totalSymbols, framesPerCondition, zValue, groupingMode)
% 按场景或场景+Eb/N0分组汇总
scenarioValues = unique(conditions.Scenario, 'stable');
result = table();
for scenario = scenarioValues.'
if strcmp(groupingMode, 'ScenarioEbNo')
ebNoValues = unique(conditions.EbNoDb( ...
conditions.Scenario == scenario), 'stable').';
else
ebNoValues = NaN;
end
for ebNoDb = ebNoValues
mask = conditions.Scenario == scenario;
if strcmp(groupingMode, 'ScenarioEbNo')
mask = mask & conditions.EbNoDb == ebNoDb;
end
numFrames = sum(mask) * framesPerCondition;
bits = sum(totalBits(mask));
symbols = sum(totalSymbols(mask));
for methodIdx = 1:numel(methodNames)
errors = sum(bitErrors(mask, methodIdx));
symbolErrorCount = sum(symbolErrors(mask, methodIdx));
frameErrorCount = sum(frameErrors(mask, methodIdx));
[berLower, berUpper] = wilson_interval( ...
errors, bits, zValue);
[ferLower, ferUpper] = wilson_interval( ...
frameErrorCount, numFrames, zValue);
plotBER = max(errors, 0.5) / bits;
row = table(scenario, ebNoDb, methodNames(methodIdx), ...
errors, bits, errors / bits, plotBER, ...
berLower, berUpper, symbolErrorCount / symbols, ...
frameErrorCount / numFrames, ferLower, ferUpper, ...
'VariableNames', {'Scenario', 'EbNoDb', 'Method', ...
'BitErrors', 'TotalBits', 'BER', 'PlotBER', ...
'BERLower', 'BERUpper', 'SER', 'FER', ...
'FERLower', 'FERUpper'});
if isempty(result)
result = row;
else
result = [result; row]; %#ok<AGROW>
end
end
end
end
end
function timingTable = build_timing_summary( ...
featureSeconds, inferenceSeconds, postprocessSeconds, ...
endToEndSeconds, zfSeconds, mmseSeconds, numSymbolsPerFrame)
% 构建时序汇总表
stageNames = [ ...
"CNN 特征提取"
"CNN 推理"
"CNN 后处理"
"CNN 端到端检测"
"ZF 检测"
"MMSE 检测"];
values = {featureSeconds, inferenceSeconds, postprocessSeconds, ...
endToEndSeconds, zfSeconds, mmseSeconds};
numStages = numel(stageNames);
meanMs = zeros(numStages, 1);
medianMs = zeros(numStages, 1);
p95Ms = zeros(numStages, 1);
p99Ms = zeros(numStages, 1);
minMs = zeros(numStages, 1);
maxMs = zeros(numStages, 1);
microsecondsPerOFDMSymbol = zeros(numStages, 1);
framesPerSecond = zeros(numStages, 1);
for stageIdx = 1:numStages
milliseconds = 1000 * values{stageIdx}(:);
meanMs(stageIdx) = mean(milliseconds);
medianMs(stageIdx) = median(milliseconds);
p95Ms(stageIdx) = percentile(milliseconds, 95);
p99Ms(stageIdx) = percentile(milliseconds, 99);
minMs(stageIdx) = min(milliseconds);
maxMs(stageIdx) = max(milliseconds);
microsecondsPerOFDMSymbol(stageIdx) = ...
1000 * meanMs(stageIdx) / numSymbolsPerFrame;
framesPerSecond(stageIdx) = 1000 / meanMs(stageIdx);
end
timingTable = table(stageNames, meanMs, medianMs, p95Ms, p99Ms, ...
minMs, maxMs, microsecondsPerOFDMSymbol, framesPerSecond, ...
'VariableNames', {'Stage', 'MeanMsPerFrame', ...
'MedianMsPerFrame', 'P95MsPerFrame', 'P99MsPerFrame', ...
'MinMsPerFrame', 'MaxMsPerFrame', ...
'MeanMicrosecondsPerOFDMSymbol', 'FramesPerSecond'});
end
function comparison = build_comparison( ...
conditionTable, summaryTable, methodNames)
% 构建CNN与ZF/MMSE的对比统计
cnnIdx = find(summaryTable.Method == "CNN", 1);
zfIdx = find(summaryTable.Method == "ZF", 1);
mmseIdx = find(summaryTable.Method == "MMSE", 1);
totalBits = summaryTable.TotalBits(cnnIdx);
cnnComparableBER = max(summaryTable.BitErrors(cnnIdx), 0.5) / totalBits;
zfComparableBER = max(summaryTable.BitErrors(zfIdx), 0.5) / totalBits;
mmseComparableBER = ...
max(summaryTable.BitErrors(mmseIdx), 0.5) / totalBits;
comparison = struct();
comparison.BERGainVsZFDb = ...
10 * log10(zfComparableBER / cnnComparableBER);
comparison.BERGainVsMMSEDb = ...
10 * log10(mmseComparableBER / cnnComparableBER);
comparison.ConditionsCNNBetterThanZF = sum( ...
conditionTable.CNNBER < conditionTable.ZFBER);
comparison.ConditionsCNNBetterThanMMSE = sum( ...
conditionTable.CNNBER < conditionTable.MMSEBER);
comparison.ConditionsCNNBestOrTied = sum(contains( ...
conditionTable.BestBERMethod, methodNames(1)));
comparison.TotalConditions = height(conditionTable);
comparison.ZeroErrorComparisonConvention = ...
"当观测到零错误时,使用0.5个错误作为比较上限";
end
function print_summary(report)
% 打印汇总结果到命令行
fprintf("\n=== 汇总结果 ===\n");
disp(report.Summary);
fprintf("CNN vs ZF: BER 增益 %.2f dB | 在 %d/%d 条件下更优\n", ...
report.Comparison.BERGainVsZFDb, ...
report.Comparison.ConditionsCNNBetterThanZF, ...
report.Comparison.TotalConditions);
fprintf("CNN vs MMSE: BER 增益 %.2f dB | 在 %d/%d 条件下更优\n", ...
report.Comparison.BERGainVsMMSEDb, ...
report.Comparison.ConditionsCNNBetterThanMMSE, ...
report.Comparison.TotalConditions);
fprintf("CNN 最优或并列最优的条件数: %d/%d\n", ...
report.Comparison.ConditionsCNNBestOrTied, ...
report.Comparison.TotalConditions);
fprintf("\n=== 响应时间 ===\n");
disp(report.Timing.Summary);
hardest = sortrows(report.Conditions, 'CNNBER', 'descend');
hardest = hardest(1:min(10, height(hardest)), ...
{'Scenario', 'KFactor', 'DopplerHz', 'EbNoDb', ...
'CNNBER', 'ZFBER', 'MMSEBER', 'CNNFER', ...
'CNNInferenceP95Ms', 'BestBERMethod'});
fprintf("\n=== CNN 最困难的10个条件 ===\n");
disp(hardest);
end
function handles = plot_report(report)
% 绘制两个图形:误码率曲线和延迟累积分布
handles = struct();
fig = figure('Name', 'CNN 深度测试 - 误码率', ...
'Position', [80 80 420 * numel(unique(report.SNRSummary.Scenario)) 430]);
handles.Figure = fig;
tiledlayout(1, numel(unique(report.SNRSummary.Scenario)), 'TileSpacing', 'compact');
for scenarioIdx = 1:numel(unique(report.SNRSummary.Scenario))
scenario = unique(report.SNRSummary.Scenario, 'stable')(scenarioIdx);
ax = nexttile;
hold(ax, 'on');
for methodIdx = 1:numel(report.MethodNames)
mask = report.SNRSummary.Scenario == scenario & ...
report.SNRSummary.Method == report.MethodNames(methodIdx);
rows = report.SNRSummary(mask, :);
semilogy(ax, rows.EbNoDb, rows.PlotBER, 'o-', ...
'LineWidth', 1.5, 'MarkerSize', 5, ...
'Color', lines(numel(report.MethodNames))(methodIdx, :), ...
'DisplayName', report.MethodNames(methodIdx));
end
title(ax, char(scenario));
xlabel(ax, 'E_b/N_0 (dB)');
ylabel(ax, '误码率 (BER)');
grid(ax, 'on');
ylim(ax, [1e-7 1]);
legend(ax, 'Location', 'southwest');
hold(ax, 'off');
end
sgtitle('CNN vs ZF vs MMSE 在同一帧上的性能对比');
% 延迟图
fig2 = figure('Name', 'CNN 延迟分布', ...
'Position', [120 120 980 420]);
handles.LatencyFigure = fig2;
tiledlayout(1, 2, 'TileSpacing', 'compact');
ax1 = nexttile;
plot_latency_cdf(ax1, 1000 * report.Timing.CNNInferenceSeconds, 'CNN 推理延迟');
ax2 = nexttile;
plot_latency_cdf(ax2, 1000 * report.Timing.CNNEndToEndSeconds, 'CNN 端到端延迟');
handles.Axes = [ax1; ax2];
end
function plot_latency_cdf(ax, valuesMs, plotTitle)
% 绘制延迟的累积分布函数 (CDF)
sortedValues = sort(valuesMs(:));
probabilities = (1:numel(sortedValues)).' / numel(sortedValues);
plot(ax, sortedValues, probabilities, 'LineWidth', 1.5);
xlabel(ax, '每帧处理时间 (ms)');
ylabel(ax, '累积概率');
title(ax, plotTitle);
grid(ax, 'on');
ylim(ax, [0 1]);
end
function [lower, upper] = wilson_interval(errors, total, zValue)
% 计算 Wilson 置信区间
totalArray = total;
if isscalar(totalArray)
totalArray = repmat(totalArray, size(errors));
else
totalArray = totalArray + zeros(size(errors));
end
proportion = errors ./ totalArray;
denominator = 1 + zValue^2 ./ totalArray;
center = (proportion + zValue^2 ./ (2 * totalArray)) ./ denominator;
margin = zValue .* sqrt( ...
proportion .* (1 - proportion) ./ totalArray + ...
zValue^2 ./ (4 * totalArray.^2)) ./ denominator;
lower = max(0, center - margin);
upper = min(1, center + margin);
end
function value = percentile(samples, percentage)
% 计算百分位数
samples = sort(samples(:));
if isempty(samples)
value = NaN;
return;
end
if isscalar(samples)
value = samples;
return;
end
position = 1 + (numel(samples) - 1) * percentage / 100;
lowerIdx = floor(position);
upperIdx = ceil(position);
interpolation = position - lowerIdx;
value = samples(lowerIdx) + interpolation * ...
(samples(upperIdx) - samples(lowerIdx));
end
function value = normal_quantile(probability)
% 标准正态分布分位数
value = -sqrt(2) * erfcinv(2 * probability);
end
function label = condition_label(condition)
% 生成条件描述字符串
if condition.Scenario == "AWGN"
label = sprintf('AWGN | Eb/No=%g dB', condition.EbNoDb);
else
label = sprintf('%s | K=%g | fd=%g Hz | Eb/No=%g dB', ...
condition.Scenario, condition.KFactor, ...
condition.DopplerHz, condition.EbNoDb);
end
end
function seed = frame_seed(baseSeed, conditionIdx, frameIdx)
% 为每个帧生成唯一的随机种子
seed = double(baseSeed) + 100000 * double(conditionIdx) + ...
double(frameIdx);
end
function scores = normalize_score_shape(scores, cnn, batchCount)
% 调整CNN输出的分数张量形状为 [M x NFFT x batch]
if ismatrix(scores)
scores = reshape(scores, size(scores, 1), size(scores, 2), 1);
end
if size(scores, 1) ~= cnn.ModOrder
error('CNN输出无效: 第一维大小不等于调制阶数 M。');
end
if size(scores, 2) == cnn.NFFT && size(scores, 3) == batchCount
return;
end
if size(scores, 2) == batchCount && size(scores, 3) == cnn.NFFT
scores = permute(scores, [1 3 2]);
return;
end
error('CNN输出形状无效: 期望 [M x NFFT x batch],实际为 [%s]。', num2str(size(scores)));
end
function validate_configuration(cfg)
% 验证输入配置参数的有效性
allowedScenarios = ["AWGN" "LOS" "NLOS"];
if isempty(cfg.Scenarios) || any(~ismember(cfg.Scenarios, allowedScenarios))
error('支持的场景: AWGN, LOS, NLOS。');
end
if numel(unique(cfg.Scenarios)) ~= numel(cfg.Scenarios)
error('场景列表不能包含重复项。');
end
validateattributes(cfg.EbNoDb, {'numeric'}, ...
{'vector', 'nonempty', 'real', 'finite'});
validateattributes(cfg.LOSKFactors, {'numeric'}, ...
{'vector', 'real', 'finite', 'nonnegative'});
validateattributes(cfg.NLOSKFactors, {'numeric'}, ...
{'vector', 'real', 'finite', 'nonnegative'});
validateattributes(cfg.DopplerHz, {'numeric'}, ...
{'vector', 'nonempty', 'real', 'finite', 'nonnegative'});
validateattributes(cfg.FramesPerCondition, {'numeric'}, ...
{'scalar', 'integer', 'positive'});
validateattributes(cfg.RandomSeed, {'numeric'}, ...
{'scalar', 'integer', 'nonnegative'});
validateattributes(cfg.WarmupIterations, {'numeric'}, ...
{'scalar', 'integer', 'nonnegative'});
validateattributes(cfg.SynchronizationThreshold, {'numeric'}, ...
{'scalar', 'real', 'finite', 'nonnegative'});
validateattributes(cfg.ConfidenceLevel, {'numeric'}, ...
{'scalar', 'real', '>', 0, '<', 1});
if any(cfg.Scenarios == "LOS") && isempty(cfg.LOSKFactors)
error('当包含LOS场景时,LOSKFactors不能为空。');
end
if any(cfg.Scenarios == "NLOS") && isempty(cfg.NLOSKFactors)
error('当包含NLOS场景时,NLOSKFactors不能为空。');
end
if ~(ischar(cfg.ModelPath) || ...
(isstring(cfg.ModelPath) && isscalar(cfg.ModelPath)))
error('ModelPath必须是标量字符串。');
end
if ~(islogical(cfg.MakePlots) && isscalar(cfg.MakePlots))
error('MakePlots必须是标量逻辑值。');
end
if ~(islogical(cfg.Verbose) && isscalar(cfg.Verbose))
error('Verbose必须是标量逻辑值。');
end
end
4 总结
-
本文基于IEEE 802.11a OFDM系统,构建了一套完整的端到端仿真平台,系统对比了基于卷积神经网络的智能均衡器与MMSE线性均衡器的性能。实验结果表明,CNN均衡器在Rician衰落信道下可获得相比MMSE更优的BER增益,且在高K因子条件下优势更为显著。本研究为深度学习技术在物理层接收机中的工程化应用提供了可复现的参考基准。
-
展望未来,深度学习在物理层接收机中的应用将朝着模型压缩与边缘部署、跨域知识迁移以及物理模型与神经网络的有机融合等方向发展。将CNN均衡器与迭代检测、信道解码等模块进行端到端联合优化,有望在复杂实际场景中实现更大的性能突破。
仿真代码可见文末VX公众号,所见即所得
更多推荐




所有评论(0)