关注我,追更更多通信仿真!

摘要

随着无线通信系统向更高频段、更复杂传播环境演进,传统信道均衡算法在性能与复杂度之间面临越来越严峻的权衡。本文基于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×104 课程学习首阶段
最终学习率 3 × 10 − 5 3 \times 10^{-5} 3×105 课程学习末阶段
训练配置 批次大小 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公众号,所见即所得

Logo

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

更多推荐