摘要

本文基于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),kD

其中 (\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 卷积,感受野可覆盖整个有效带宽。整体结构如下:

  1. Stem 层:1×1 卷积,将输入通道从 10 提升至 64,后接 ReLU。
  2. 5 个残差块:每个块包含两层空洞卷积(卷积核大小 9,空洞率依次为 1, 2, 4, 8, 1),每层后接 ReLU,并通过跳跃连接将输入加到输出上(残差结构)。空洞卷积在不增加参数的前提下指数级扩大感受野,使网络能够利用远端子载波的相关性。
  3. 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=kDm=1Mtm(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公众号

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐