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

摘要

正交频分复用(OFDM)技术在现代无线通信系统中占据核心地位,但其高峰均比(PAPR)特性使发射信号对功率放大器(PA)的非线性失真极为敏感。传统线性均衡器(如MMSE)基于线性信道模型设计,无法补偿PA非线性引起的跨子载波互调干扰,导致高信噪比区域出现误码率平台效应。本文基于IEEE 802.11a标准构建了完整的OFDM物理层仿真链路,系统评估了卷积神经网络(CNN)均衡器在Rapp固态功率放大器(SSPA)非线性失真下的误码率性能。仿真结果表明,在严重非线性条件下,CNN均衡器相比传统MMSE均衡器更优下ber性能,验证了深度学习在非线性信道均衡中的有效性。

1 背景意义

1.1 OFDM技术的核心地位与固有挑战

正交频分复用(OFDM)以其卓越的抗多径衰落能力和高频谱效率,已成为现代无线通信物理层的基石。从IEEE 802.11系列(Wi-Fi)到3GPP LTE/5G NR,从数字视频广播(DVB)到电力线通信,OFDM的应用遍及无线通信的各个领域。其核心优势在于通过循环前缀(CP)和频域均衡,将宽带频率选择性信道转化为一组并行的平坦衰落子信道,极大简化了接收机设计。

然而,OFDM系统在工程实现中面临一个固有挑战:高峰均比(Peak-to-Average Power Ratio, PAPR)。当多个子载波的信号同相叠加时,瞬时功率可比平均功率高出10 dB以上。这一特性对发射机射频前端的功率放大器(Power Amplifier, PA)提出了极其严苛的线性度要求。

1.2 功率放大器非线性失真的物理根源与实际影响

功率放大器是无线发射机中最关键的器件,其作用是将基带信号放大至足够功率以克服路径损耗。从物理层面看,PA由有源器件(如LDMOS、GaN HEMT)构成,其输入-输出传输特性天然具有非线性。当输入信号幅度较小时,PA工作在线性区;随着输入幅度增大,PA进入压缩区,增益开始下降;进一步增大至饱和区后,输出幅度趋于恒定。

PA的非线性对OFDM系统造成双重危害:

  • 带内失真:非线性互调产物落在信号带宽内,导致星座点偏移和弥散,误差矢量幅度(EVM)恶化,误码率(BER)显著升高。对于高阶调制(如64-QAM),这一影响尤为致命。

  • 带外辐射:非线性分量扩展至信号带宽之外,产生频谱再生(Spectral Regrowth),可能造成邻道干扰,违反无线电管理法规。

工程中通常用输入回退(Input Back-Off, IBO)来量化PA的工作点:
在这里插入图片描述

1.3 传统方案的局限性

为缓解非线性失真,工程中常采用两种方案:

  • 功率回退:增大IBO牺牲效率换取线性度。在电池供电的终端设备中,效率损失直接影响续航时间。

  • 数字预失真(Digital Pre-Distortion, DPD):在数字域对PA非线性进行预补偿。然而,DPD需要复杂的反馈链路和精确的PA建模,实现成本较高,且对PA特性漂移(温度、老化、频率等)敏感。

在这里插入图片描述

1.4 深度学习引入的意义

近年来,深度学习为物理层通信带来了新的可能性。与依赖精确数学模型的方法不同,深度神经网络以数据驱动的方式从样本中学习复杂映射。在非线性补偿场景中:

  • 无需先验建模:网络可自行学习PA非线性特征

  • 端到端优化:可直接优化与BER相关的代理损失

  • 泛化能力:在训练覆盖范围内可适应不同工作条件

1.5 本文意义

本文基于IEEE 802.11a标准构建了完整的OFDM物理层仿真链路,系统评估了卷积神经网络(CNN)均衡器在Rapp固态功率放大器(SSPA)非线性失真下的误码率性能。

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 功率放大器非线性模型:Rapp AM/AM

在这里插入图片描述

在这里插入图片描述

2.3 线性均衡器

在这里插入图片描述

2.3.1 迫零(ZF)均衡器

在这里插入图片描述

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

在这里插入图片描述

2.4 CNN均衡器

传统线性均衡器(ZF/MMSE)基于频域对角信道模型设计,各子载波独立处理。当PA非线性存在时,频域信道矩阵不再是对角的——非线性产生跨子载波的互调干扰,各子载波间产生耦合。MMSE均衡器仅利用单子载波信息Hk,无法补偿这种耦合,这正是其在非线性信道中性能受限的根本原因。

卷积神经网络(CNN)均衡器采用数据驱动的端到端学习范式,直接从接收的时频网格中学习到发送符号的非线性映射关系:
在这里插入图片描述

2.4.1 输入特征

针对HPA非线性场景,输入特征设计的关键原则是不做任何功率归一化,使网络能够从绝对信号幅度中感知IBO水平(非线性程度的直接指标)。对第k个子载波提取10维特征向量:
在这里插入图片描述

2.4.2 网络结构

在这里插入图片描述

2.4.3 训练和推理过程

在这里插入图片描述

  • 训练采用课程学习策略:从高IBO(线性区)样本开始,逐步引入低IBO(非线性区)样本。具体分为5个阶段,IBO范围从0–3 dB逐步扩展至−5–25 dB,学习率从0.001递减至0.00001。最后阶段引入Hard Case Mining,以60%概率采样低IBO(−5至3 dB)和高SNR(8–35 dB)的困难组合,强化网络在最具挑战工况下的泛化能力。

推理阶段取Softmax输出中最大概率对应的类别作为判决:
在这里插入图片描述
再通过Gray映射转换为比特序列,与发送真值比较统计BER。

3 仿真流程设计

3.1 系统框架

整个仿真链路严格遵循 IEEE 802.11a 物理层标准,涵盖发射机、非线性信道、接收机同步与均衡检测三大模块。整体架构如下图所示
在这里插入图片描述

  • 发射端:首先生成随机的 QPSK 调制比特序列(每子载波 2 比特),经 Gray 映射后形成频域复符号。在 52 个有效载波中,48 个为数据载波,4 个为导频载波,导频符号携固定极性序列用于辅助信道估计和相位跟踪。随后执行 IFFT 将频域符号转换至时域,并插入 16 个采样点的循环前缀。最后在数据符号前级联由 L-STF(短训练序列)和 L-LTF(长训练序列)构成的完整前导码,形成完整的物理层帧。

  • 非线性信道:发射信号首先经过 Rapp 固态功率放大器(SSPA)模型的 AM/AM 非线性变换,其非线性程度由输入回退(IBO)参数控制。IBO 越低,信号越频繁进入饱和区,非线性失真越严重。随后叠加 AWGN。采用公共随机数控制确保同一帧在不同均衡器之间保持完全相同的信道实现和噪声样本,这是保证对比公平性的关键设计。

  • 接收端:通过前导码相关检测实现帧同步,定位帧起始位置;去除 CP 后执行 FFT 将信号转换回频域;利用 4 个导频子载波进行 MMSE 信道估计,为线性均衡器提供信道状态信息。

  • 均衡检测:两种均衡器对完全相同的一组接收帧依次处理。MMSE 基于导频信道估计值进行线性滤波,在干扰消除与噪声放大之间寻求最优平衡;CNN 则先提取多通道特征张量(保留全频带 I/Q、导频 MMSE 均衡结果、信道增益、导频残留误差等),再执行神经网络前向推理,输出 QPSK 各符号类别的概率分布。两种方法的输出均为 QPSK 符号序列,经 Gray 解调后恢复比特,与发射端真值逐位比较统计 BER。

3.2 仿真参数

3.2.1 802.11a OFDM 参数
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) 纯粹评估均衡器性能
3.2.2 非线性信道
非线性信道参数名称 符号/取值 说明
非线性模型 Rapp AM/AM 固态功率放大器模型
平滑度因子 p = 2 p = 2 p=2 固定值, p p p 越大越接近限幅器
饱和电压 V s a t = 1 V_{sat} = 1 Vsat=1 归一化饱和电平
输入回退范围 IBO = − 10 , − 5 , 0 , 5 -10, -5, 0, 5 10,5,0,5 dB 覆盖过驱动到线性区
噪声模型 AWGN 叠加于 HPA 输出后
3.2.3 CNN超参数
CNN参数 参数名称 取值 说明
输入层 输入通道数 C = 10 C = 10 C=10 见 2.5.2 节特征构造
输入维度 C × 64 C \times 64 C×64 全频带(含 DC/Guard)
卷积层 卷积核大小 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 概率归一化
残差结构 残差块数量 5 含跳跃连接
损失函数 类型 交叉熵 + 软符号 MSE 混合损失
软符号 MSE 权重 λ = 0.5 → 1.0 \lambda = 0.5 \to 1.0 λ=0.51.0 随阶段递增
优化器 类型 Adam 自适应矩估计
初始学习率 1.0 × 10 − 3 1.0 \times 10^{-3} 1.0×103 课程学习首阶段
最终学习率 7.0 × 10 − 5 7.0 \times 10^{-5} 7.0×105 课程学习末阶段
训练配置 批次大小 256 每批训练样本数
每轮样本数 6.0 × 10 3 → 2.4 × 10 4 6.0 \times 10^3 \to 2.4 \times 10^4 6.0×1032.4×104 逐阶段递增
训练阶段数 5 课程学习

3.3 仿真图分析

在这里插入图片描述

可以看到:

  • 非线性越严重,CNN优势越明显。IBO = -10 dB和-5 dB时CNN在所有19个SNR点均优于MMSE;IBO = -5 dB降幅最大,优于IBO = -10 dB。原因在于IBO = -10 dB信号被过度压缩,信息已不可逆丢失;IBO = -5 dB非线性仍显著但保留了足够结构信息供CNN学习补偿。

  • 低SNR(0 dB)时CNN和MMSE差异较小,噪声主导

  • 综上,CNN有效补偿了MMSE无法处理的跨子载波互调干扰。

部分代码:

classdef HPA_CNN < handle
    properties
        NFFT (1,1) double = 64
        ModOrder (1,1) double = 4
        NumInputChannels (1,1) double = 10
        MiniBatchSize (1,1) double = 64
        MaxEpochs (1,1) double = 50
        InitialLearnRate (1,1) double = 1e-3
        SymbolMSEWeight (1,1) double = 0.5
        Constellation
        Net
        TrainingInfo
    end

    methods
        function obj = HPA_CNN(varargin)
            % 输入参数必须为名称/值对
            if mod(nargin, 2) ~= 0
                error('参数必须是名称/值对。');
            end

            for k = 1:2:nargin
                obj.(varargin{k}) = varargin{k+1};
            end

            obj.validateModOrder();
            if obj.NumInputChannels < 2
                error('CNN至少需要两个I/Q输入通道。');
            end
            if obj.SymbolMSEWeight < 0
                error('SymbolMSEWeight不能为负数。');
            end

            obj.Constellation = obj.buildConstellation();
            obj.Net = dlnetwork(obj.buildLayerGraph());
        end

        function lgraph = buildLayerGraph(obj)
            % 构建网络层图
            lgraph = layerGraph([
                sequenceInputLayer(obj.NumInputChannels, ...
                    'Name', 'input', ...
                    'Normalization', 'none', ...
                    'MinLength', obj.NFFT)
                convolution1dLayer(1, 64, 'Padding', 'same', 'Name', 'stem')
                reluLayer('Name', 'stem_relu')
                ]);

            previousLayer = 'stem_relu';
            dilations = [1 2 4 8 1];
            for b = 1:numel(dilations)
                blockLayers = [
                    convolution1dLayer(9, 64, 'Padding', 'same', ...
                        'DilationFactor', dilations(b), ...
                        'Name', sprintf('res%d_conv1', b))
                    reluLayer('Name', sprintf('res%d_relu1', b))
                    convolution1dLayer(9, 64, 'Padding', 'same', ...
                        'DilationFactor', dilations(b), ...
                        'Name', sprintf('res%d_conv2', b))
                    additionLayer(2, 'Name', sprintf('res%d_add', b))
                    reluLayer('Name', sprintf('res%d_relu_out', b))
                    ];

                lgraph = addLayers(lgraph, blockLayers);
                lgraph = connectLayers(lgraph, previousLayer, sprintf('res%d_conv1', b));
                lgraph = connectLayers(lgraph, previousLayer, sprintf('res%d_add/in2', b));
                previousLayer = sprintf('res%d_relu_out', b);
            end

            headLayers = [
                convolution1dLayer(1, 48, 'Padding', 'same', 'Name', 'carrier_mixer')
                reluLayer('Name', 'head_relu')
                convolution1dLayer(1, obj.ModOrder, 'Padding', 'same', 'Name', 'class_scores')
                softmaxLayer('Name', 'softmax')
                ];

            lgraph = addLayers(lgraph, headLayers);
            lgraph = connectLayers(lgraph, previousLayer, 'carrier_mixer');
        end

        function X = normalizeInputShape(obj, X)
            % 输入CNN: [NumInputChannels x NFFT x batch]
            if ~isreal(X)
                error('无效输入:X必须包含分离的实部I/Q特征。');
            end

            if ismatrix(X)
                if size(X, 1) ~= obj.NumInputChannels || size(X, 2) ~= obj.NFFT
                    error('无效的2D输入。期望 [NumInputChannels x NFFT]。');
                end
                X = reshape(X, obj.NumInputChannels, obj.NFFT, 1);
            elseif ndims(X) ~= 3
                error('无效输入。期望 [NumInputChannels x NFFT x batch]。');
            end

            if size(X, 1) ~= obj.NumInputChannels || size(X, 2) ~= obj.NFFT
                error('无效输入。期望 [NumInputChannels x NFFT x batch]。');
            end
        end

        function Xn = prepareInputBatch(obj, X, carrierIdx)
            % 准备输入批次
            X = obj.normalizeInputShape(X);

            if nargin >= 3 && ~isempty(carrierIdx)
                carrierIdx = double(carrierIdx(:)).';
                if any(carrierIdx < 1) || any(carrierIdx > obj.NFFT)
                    error('carrierIdx超出NFFT=%d的范围。', obj.NFFT);
                end
            end

            if any(~isfinite(X(:)))
                error('CNN特征包含非有限值。');
            end

            % 无AGC、标准化或裁剪:绝对电平是识别IBO所需信息的一部分。
            Xn = single(X);
        end

        function Xn = normalizeInputBatch(obj, X, carrierIdx)
            % 向后兼容别名。不执行任何归一化。
            if nargin < 3
                carrierIdx = [];
            end
            Xn = obj.prepareInputBatch(X, carrierIdx);
        end

        function T = normalizeTargetShape(obj, T)
            % 目标:one-hot [M x NFFT x batch] 或类别 [NFFT x batch]
            if ismatrix(T)
                if size(T, 1) == obj.ModOrder && size(T, 2) == obj.NFFT && obj.looksLikeOneHot(T)
                    T = reshape(single(T), obj.ModOrder, obj.NFFT, 1);
                elseif size(T, 1) == obj.NFFT
                    T = obj.classesToOneHot(T);
                else
                    error('无效的2D目标。期望 [NFFT x batch][M x NFFT]。');
                end
            elseif ndims(T) ~= 3
                error('无效目标。期望 [M x NFFT x batch]。');
            end

            if size(T, 1) ~= obj.ModOrder || size(T, 2) ~= obj.NFFT
                error('无效的one-hot目标。期望 [M x NFFT x batch]。');
            end

            obj.validateOneHot(T);
            T = single(T);
        end

        function oneHot = classesToOneHot(obj, classes)
            % classes: [NFFT x batch],值为整数 1..M
            if isvector(classes)
                classes = classes(:);
            end
            if ~ismatrix(classes)
                error('无效的类别标签。期望 [NFFT x batch]。');
            end

            classes = obj.validateClasses(classes);
            [nfft, nframes] = size(classes);
            oneHot = zeros(obj.ModOrder, nfft, nframes, 'single');

            classIdx = classes(:);
            carrierIdx = repmat((1:nfft).', nframes, 1);
            frameIdx = reshape(repmat(1:nframes, nfft, 1), [], 1);
            linearIdx = sub2ind([obj.ModOrder, nfft, nframes], classIdx, carrierIdx, frameIdx);
            oneHot(linearIdx) = 1;
        end

        function classes = oneHotToClasses(obj, oneHot)
            % oneHot: [M x NFFT x batch][M x NFFT]
            if ismatrix(oneHot)
                if size(oneHot, 1) ~= obj.ModOrder
                    error('无效的one-hot。第一维必须是M。');
                end
                [~, idx] = max(oneHot, [], 1);
                classes = reshape(double(idx), size(oneHot, 2), 1);
            elseif ndims(oneHot) == 3
                if size(oneHot, 1) ~= obj.ModOrder
                    error('无效的one-hot。第一维必须是M。');
                end
                [~, idx] = max(oneHot, [], 1);
                classes = reshape(double(idx), size(oneHot, 2), size(oneHot, 3));
            else
                error('无效的one-hot。期望 [M x NFFT x batch]。');
            end
        end

        function symbols = classesToSymbols(obj, classes)
            % 内部类别为1..M;qammod使用索引0..M-1
            classes = obj.validateClasses(classes);
            qamIntegers = classes(:) - 1;
            symbols = qammod(qamIntegers, obj.ModOrder, 'gray', 'UnitAveragePower', true);
            symbols = reshape(symbols, size(classes));
        end

        function symbols = scoresToSymbols(obj, scores)
            % scores: [M x NFFT x batch],网络softmax输出
            if ismatrix(scores)
                scores = reshape(scores, size(scores, 1), size(scores, 2), 1);
            end
            if ndims(scores) ~= 3 || size(scores, 1) ~= obj.ModOrder || ...
                    size(scores, 2) ~= obj.NFFT
                error('无效的得分。期望 [M x NFFT x batch]。');
            end
            symbols = sum(scores .* reshape(obj.Constellation, [], 1, 1), 1);
            symbols = reshape(symbols, obj.NFFT, size(scores, 3));
        end

        function classes = symbolsToClasses(obj, symbols)
            qamIntegers = qamdemod(symbols(:), obj.ModOrder, 'gray', ...
                'UnitAveragePower', true, 'OutputType', 'integer');
            classes = reshape(double(qamIntegers) + 1, size(symbols));
        end

        function bits = classesToBits(obj, classes)
            symbols = obj.classesToSymbols(classes);
            bits = obj.symbolsToBits(symbols);
        end

        function classes = bitsToClasses(obj, bits)
            symbols = obj.bitsToSymbols(bits);
            classes = obj.symbolsToClasses(symbols);
        end

        function symbols = bitsToSymbols(obj, bits)
            [bits, outSize] = obj.validateBits(bits);
            symbols = qammod(double(bits(:)), obj.ModOrder, 'gray', ...
                'InputType', 'bit', 'UnitAveragePower', true);
            symbols = reshape(symbols, outSize);
        end

        function bits = symbolsToBits(obj, symbols)
            bitsColumn = qamdemod(symbols(:), obj.ModOrder, 'gray', ...
                'UnitAveragePower', true, 'OutputType', 'bit');
            bits = reshape(uint8(bitsColumn(:)), [obj.bitsPerSymbol(), size(symbols)]);
        end

        function ber = computeBER(obj, trueClasses, predClasses)
            trueClasses = obj.validateClasses(trueClasses);
            predClasses = obj.validateClasses(predClasses);
            if ~isequal(size(trueClasses), size(predClasses))
                error('trueClasses和predClasses必须具有相同维度。');
            end

            trueBits = obj.classesToBits(trueClasses);
            predBits = obj.classesToBits(predClasses);
            ber = mean(trueBits(:) ~= predBits(:));
        end

        function evm = computeEVM(~, trueSymbols, predSymbols)
            num = sum(abs(trueSymbols - predSymbols).^2, 'all');
            den = sum(abs(trueSymbols).^2, 'all') + eps;
            evm = sqrt(num / den);
        end

        function mse = computeMSEOneHot(obj, trueClasses, predClasses)
            trueOneHot = obj.classesToOneHot(trueClasses);
            predOneHot = obj.classesToOneHot(predClasses);
            mse = mean((trueOneHot(:) - predOneHot(:)).^2);
        end
    end

    methods (Static)
        function version = architectureVersion()
            version = "carrier_resnet_v1_rapp_soft_symbol_loss";
        end

        function Y = reshapeToClassTime(X, M, NFFT)
            Y = reshape(X, [M, NFFT, size(X, 2)]);
        end

        function loss = crossEntropyLoss(dlY, dlT, carrierIdx)
            if nargin >= 3 && ~isempty(carrierIdx)
                carrierIdx = double(carrierIdx(:)).';
                dlY = HPA_CNN.selectCarrierDimension(dlY, carrierIdx);
                dlT = HPA_CNN.selectCarrierDimension(dlT, carrierIdx);
            end

            carrierDim = HPA_CNN.dimensionByLabel(dlT, 'T', 2);
            batchDim = HPA_CNN.dimensionByLabel(dlT, 'B', 3);
            lossTerms = -dlT .* log(dlY + single(1e-8));
            loss = sum(lossTerms, 'all') / (size(dlT, carrierDim) * size(dlT, batchDim));
        end

        function [loss, gradients, crossEntropy, symbolMSE] = modelGradients( ...
                net, dlX, dlT, carrierIdx, constellation, symbolMSEWeight)
            if nargin < 4
                carrierIdx = [];
            end
            if nargin < 5 || isempty(constellation)
                error('必须将星座传递给modelGradients。');
            end
            if nargin < 6 || isempty(symbolMSEWeight)
                symbolMSEWeight = 0.5;
            end

            dlY = forward(net, dlX);
            crossEntropy = HPA_CNN.crossEntropyLoss(dlY, dlT, carrierIdx);
            symbolMSE = HPA_CNN.softSymbolMSE( ...
                dlY, dlT, carrierIdx, constellation);
            loss = crossEntropy + symbolMSEWeight * symbolMSE;
            gradients = dlgradient(loss, net.Learnables);
        end

        function loss = softSymbolMSE(dlY, dlT, carrierIdx, constellation)
            if nargin >= 3 && ~isempty(carrierIdx)
                dlY = HPA_CNN.selectCarrierDimension(dlY, carrierIdx);
                dlT = HPA_CNN.selectCarrierDimension(dlT, carrierIdx);
            end

            constellation = constellation(:);
            cReal = reshape(single(real(constellation)), [], 1, 1);
            cImag = reshape(single(imag(constellation)), [], 1, 1);
            predReal = sum(dlY .* cReal, 1);
            predImag = sum(dlY .* cImag, 1);
            trueReal = sum(dlT .* cReal, 1);
            trueImag = sum(dlT .* cImag, 1);
            loss = mean((predReal - trueReal).^2 + ...
                (predImag - trueImag).^2, 'all');
        end

        function dlX = selectCarrierDimension(dlX, carrierIdx)
            carrierDim = HPA_CNN.dimensionByLabel(dlX, 'T', 2);
            subs = repmat({':'}, 1, max(ndims(stripdims(dlX)), carrierDim));
            subs{carrierDim} = carrierIdx;
            dlX = dlX(subs{:});
        end

        function dim = dimensionByLabel(dlX, label, defaultDim)
            dim = defaultDim;
            try
                fmt = char(dims(dlX));
                idx = find(fmt == label, 1);
                if ~isempty(idx)
                    dim = idx;
                end
            catch
            end
        end
    end

    methods (Access = private)
        function constellation = buildConstellation(obj)
            qamIntegers = (0:obj.ModOrder-1).';
            constellation = qammod(qamIntegers, obj.ModOrder, 'gray', ...
                'UnitAveragePower', true);
            constellation = constellation(:);
        end

        function validateModOrder(obj)
            k = log2(obj.ModOrder);
            if obj.ModOrder < 2 || abs(k - round(k)) > eps
                error('ModOrder必须是2的幂。');
            end
        end

        function k = bitsPerSymbol(obj)
            k = round(log2(obj.ModOrder));
        end

        function classes = validateClasses(obj, classes)
            if isempty(classes)
                error('类别标签为空。');
            end
            if any(~isfinite(double(classes(:))))
                error('类别非有限值。');
            end
            if any(abs(double(classes(:)) - round(double(classes(:)))) > eps)
                error('类别必须是整数。');
            end
            if any(classes(:) < 1) || any(classes(:) > obj.ModOrder)
                error('类别超出范围。必须在1和M之间。');
            end
            classes = double(classes);
        end

        function [bits, outSize] = validateBits(obj, bits)
            if isempty(bits)
                error('比特为空。');
            end

            k = obj.bitsPerSymbol();
            if isvector(bits) && size(bits, 1) ~= k
                if mod(numel(bits), k) ~= 0
                    error('比特数不能被log2(M)整除。');
                end
                bits = reshape(bits, k, []);
            end

            if size(bits, 1) ~= k
                error('无效的比特。期望 [log2(M) x NFFT x batch]。');
            end
            if any(bits(:) ~= 0 & bits(:) ~= 1)
                error('比特必须是0或1。');
            end

            bits = uint8(bits);
            sz = size(bits);
            if ismatrix(bits)
                outSize = [sz(2), 1];
            else
                outSize = sz(2:end);
            end
        end

        function validateOneHot(~, oneHot)
            if any(oneHot(:) < -1e-6) || any(oneHot(:) > 1 + 1e-6)
                error('无效的one-hot目标:值超出[0,1]。');
            end

            sums = sum(oneHot, 1);
            if any(abs(sums(:) - 1) > 1e-4)
                error('无效的one-hot目标:每个载波的和必须为1。');
            end
        end

        function tf = looksLikeOneHot(~, T)
            tf = all(abs(T(:) - round(T(:))) < 1e-6) && ...
                all(T(:) >= 0) && all(T(:) <= 1) && ...
                all(abs(sum(T, 1) - 1) < 1e-4);
        end
    end
end

4 总结

  • 本文针对OFDM系统在功率放大器非线性失真条件下的均衡问题,基于IEEE 802.11a标准构建了完整的物理层仿真链路,系统评估了基于一维残差卷积神经网络的智能均衡器的性能,并与传统MMSE线性均衡器进行了全面的对比分析。研究工作涵盖了发射机建模、Rapp固态功率放大器非线性失真模拟、接收机同步与信道估计、CNN特征提取与网络推理、以及基于课程学习的训练策略等关键环节。

  • 在IBO受限的系统中——如低功耗物联网终端、手持设备、卫星通信终端等场景——当数字预失真(DPD)因成本、功耗或反馈链路限制而难以部署时,CNN均衡器可作为纯软件升级方案部署于接收端基带处理器,在不改变发射机硬件的前提下改善链路性能。

  • 未来可以考虑更广泛的非线性模型如含记忆效应的PA(如Wiener、Hammerstein模型)、AM/PM失真、以及功放随温度/频率/老化漂移的鲁棒性评估。以及考虑迁移学习与在线自适应:部署后的PA特性可能因环境变化而漂移。研究基于少量在线数据的迁移学习或自适应微调策略,可使CNN均衡器在部署后持续适应变化的非线性特性。

仿真代码可见文末VX公众号,所见即所得

更多推荐