第【39】期--基于IEEE 802.11a的深度学习接收端均衡器设计 --matlab完整代码
文章目录
摘要
本文基于IEEE 802.11a标准,构建了一个端到端的OFDM仿真平台,设计并实现了一种轻量化的1D-CNN均衡器。通过大量蒙特卡洛仿真,在AWGN、Rician和NLOS多径信道下系统评估了CNN与传统MMSE均衡器的性能差异。
1. 背景与意义
在无线通信系统中,正交频分复用(OFDM)凭借其抗多径衰落和频谱效率高的特性,成为4G/5G/Wi-Fi等现代标准的核心物理层技术。然而,OFDM接收机中信道均衡环节的性能直接决定了系统误码率(BER)。传统均衡器(如ZF和MMSE)基于线性滤波理论,在频率选择性衰落信道下存在明显的噪声放大或残留干扰问题,尤其在低信噪比或深衰落场景中性能急剧恶化。
近年来,深度学习(DL)在物理层通信中的应用如火如荼,其中基于卷积神经网络(CNN)的数据驱动均衡器被证明能够隐式学习信道逆映射,甚至捕捉非线性失真,从而超越线性均衡器。但大多数研究仅停留在理论仿真,缺乏对完整物理层帧结构、信道模型和工程实现的系统化验证。
本文基于IEEE 802.11a标准的物理层帧结构,搭建了一个端到端的OFDM仿真平台,并训练了一个面向子载波的轻量化1D-CNN均衡器。通过大量蒙特卡洛仿真,在AWGN、Rician(不同K因子)和NLOS多径场景下对比了CNN与MMSE的BER性能,同时评估了CNN的推理延迟,旨在为深度学习在真实通信系统中的应用提供可复现的实验参考。
2. 理论基础
2.1 OFDM 系统模型(IEEE 802.11a)
本仿真采用 IEEE 802.11a 物理层帧结构,核心参数如下:
| 参数 | 值 |
|---|---|
| FFT 点数 | 64 |
| 循环前缀长度 | 16 个样点 |
| 有效子载波数 | 52(48 个数据 + 4 个导频) |
| 调制方式 | QPSK(4‑QAM),无信道编码 |
| 每帧数据符号数 | 50 |
| 帧结构 | L‑STF + L‑LTF + 50 个 QPSK 符号 |
发送端将 QPSK 符号映射到 48 个数据子载波上,并插入 4 个导频子载波(位置为 ‑21, ‑7, 7, 21)。经过 IFFT 和加循环前缀后形成时域信号。接收端首先利用 L‑STF 进行粗同步,再利用 L‑LTF 估计每个子载波上的信道频率响应 (H(k)) 和噪声方差 (\sigma^2)。设接收到的频域数据子载波为:
Y ( k ) = H ( k ) ⋅ X ( k ) + N ( k ) , k ∈ D Y(k) = H(k) \cdot X(k) + N(k), \quad k \in \mathcal{D} Y(k)=H(k)⋅X(k)+N(k),k∈D
其中 (\mathcal{D}) 为数据子载波索引集合,(X(k)) 为发送的 QPSK 符号,(N(k)) 为复高斯噪声,方差为 (\sigma^2)。
传统线性均衡器:
-
迫零(ZF)均衡器:直接求信道逆,完全消除码间干扰,但会放大噪声。
X ^ ZF ( k ) = Y ( k ) H ( k ) \hat{X}_{\text{ZF}}(k) = \frac{Y(k)}{H(k)} X^ZF(k)=H(k)Y(k) -
最小均方误差(MMSE)均衡器:在干扰消除与噪声放大之间折中,使均方误差最小。
X ^ MMSE ( k ) = H ∗ ( k ) ∣ H ( k ) ∣ 2 + σ 2 ⋅ Y ( k ) \hat{X}_{\text{MMSE}}(k) = \frac{H^*(k)}{|H(k)|^2 + \sigma^2} \cdot Y(k) X^MMSE(k)=∣H(k)∣2+σ2H∗(k)⋅Y(k)
2.2 CNN 均衡器设计
本文提出的 CNN 均衡器不直接求解信道逆,而是通过学习的方式从接收信号、信道估计等特征中直接映射到发送符号的类别概率。其核心思路是将每个 OFDM 符号视为一张“频域图像”,利用 1D 卷积神经网络沿子载波维度进行特征提取与分类。
2.2.1 输入特征工程
为充分利用先验信息,我们为每个 OFDM 符号构建了一个 10 通道 的输入张量,维度为 (10 \times 64 \times B),其中 64 为 FFT 长度,(B) 为批大小(同时处理的符号数)。各通道含义如下:
| 通道 | 特征 | 说明 |
|---|---|---|
| 1, 2 | 接收信号平均(I/Q) | 去除导频调制后的活动子载波平均,取实部/虚部 |
| 3, 4 | ZF 均衡符号(I/Q) | 使用 L‑LTF 估计的信道进行 ZF 均衡后的实部/虚部 |
| 5, 6 | 信道估计(I/Q) | 频域信道响应 (H(k)) 的实部/虚部 |
| 7 | 信道幅值 | ( |
| 8 | 导频掩码 | 导频位置为 1,其他为 0 |
| 9 | 有效子载波掩码 | 活动子载波位置为 1,DC/保护带为 0 |
| 10 | 归一化子载波坐标 | 子载波索引除以最大绝对值,使值域在 ([-1, 1]) |
特征构建时,先对每个符号提取活动子载波上的接收符号,结合导频符号和正则化后的信道估计,再组装成张量,实现并行处理。
2.2.2 网络架构
CNN 采用 残差空洞卷积 结构,沿频域方向进行 1D 卷积,感受野可覆盖整个有效带宽。整体结构如下:
- Stem 层:1×1 卷积,将输入通道从 10 提升至 64,后接 ReLU。
- 5 个残差块:每个块包含两层空洞卷积(卷积核大小 9,空洞率依次为 1, 2, 4, 8, 1),每层后接 ReLU,并通过跳跃连接将输入加到输出上(残差结构)。空洞卷积在不增加参数的前提下指数级扩大感受野,使网络能够利用远端子载波的相关性。
- Head 层:两个 1×1 卷积,分别将通道降至 48 和调制阶数 (M=4),最后通过 Softmax 输出每个子载波上 4 个 QPSK 星座点的后验概率。
网络总参数量约 6.8 万,属于轻量级模型。
2.2.3 训练策略
- 动态数据生成:训练样本实时生成,每批数据来自随机生成的 802.11a 帧,经过随机化的 Rician 信道(SNR、K 因子、多径、多普勒均随机),并加入随机相位旋转以增强泛化性。
- 课程学习:训练分为多个宏观阶段,每个阶段对 SNR、K 因子、多径剖面采用不同采样区间,逐步从高 SNR 简单信道过渡到低 SNR 复杂信道,使网络循序渐进地学习。
- 损失函数与优化:采用交叉熵损失:
L = − ∑ k ∈ D ∑ m = 1 M t m ( k ) log ( p m ( k ) ) \mathcal{L} = -\sum_{k \in \mathcal{D}} \sum_{m=1}^{M} t_m(k) \log\left(p_m(k)\right) L=−k∈D∑m=1∑Mtm(k)log(pm(k))
其中 (t_m(k)) 为真实类别的 one‑hot 指示,(p_m(k)) 为网络输出的概率。优化器为 Adam,学习率随阶段递减((8\times10^{-4} \to 3\times10^{-5}))。 - 输入归一化:对每个 batch,按 I/Q 通道对独立进行能量归一化,使网络对信号绝对幅度不敏感。
2.2.4 推理与延迟
测试时,对每个接收帧提取所有符号的特征,归一化后输入网络,得到 (M \times 64 \times N_{\text{sym}}) 的分数张量,取最大值的类索引作为判决结果,再映射为比特与真实比特比较。
为提升吞吐量,采用批量推理(将多个帧的符号拼为一个大 batch)。在 Intel Xeon CPU 上,单帧(50 个符号)端到端平均耗时约 2.55 ms(特征提取 0.85 ms,前向 1.20 ms,后处理 0.50 ms),远小于 802.11a 帧周期(200 μs),具备实时部署潜力。
3. 仿真系统设计
3.1 仿真平台与参数
- 帧结构:802.11a Legacy,QPSK(无FEC),每帧50个数据符号
- 信道模型:
- AWGN:仅加性白噪声
- Rician(LOS):两径(0 ns, 30 ns),增益 [0, -10] dB,K因子可变(0~40 dB)
- NLOS:三径(0, 60, 140 ns),增益 [0, -2, -6] dB,K因子≈0(近似Rayleigh)
- 均衡器对比:CNN、MMSE
- 蒙特卡洛准则:每个SNR点至少运行40帧,达到目标误比特数(通常250~400)后提前终止,保证统计可靠性。
3.2 测试场景设计
3.3 CNN推理实现细节
在测试中,为减少反复调用网络的overhead,采用批量推理:将多个帧的OFDM符号合并为一个大batch(例如每次处理4帧共200个符号),一次性送入网络。输入经过归一化(对IQ通道进行能量缩放),输出分数张量形状为 [M × N_FFT × batch],映射回原始子载波顺序后与真实QPSK符号比对,计算误比特率。
4. 仿真结果与分析
4.2 Rician K因子对CNN性能的影响


为深入探究CNN的泛化边界,在固定两径信道下扫描K因子,可以看到大部分下均优于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
5. 结论与展望
本文通过完整的端到端仿真,系统评估了CNN均衡器在IEEE 802.11a OFDM系统中的性能。未来的工作可从以下方向展开:
- 引入在线自适应,利用接收数据微调网络,应对时变信道。
- 联合信道估计与均衡,将CNN嵌入接收机前端,实现端到端信号检测。
- 搭建FPGA或ARM原型,验证实际无线环境中的吞吐量和鲁棒性。
源代码 与仿真图表所见即所得,完整代码获取方式请见文末vx公众号
更多推荐




所有评论(0)