简介:这份PDF文献面向通信工程、信号处理方向的研究生与科研人员,聚焦深度学习在OFDM系统信号检测中的应用,帮助读者理解如何用深度神经网络替代传统检测模块。全文围绕一套完整检测框架展开:先用迫零均衡器重构DNN输入,再在离线训练中引入预训练阶段,以导频符号和数据符号作为训练数据提供良好初始参数,最后加载最优参数完成在线检测。文中给出信噪比26dB下无预训练、无ZF均衡器分别损失2dB和7dB的对比结果,并讨论导频减少、无循环前缀及不同信道参数下的误码率表现,还涉及迁移学习与数据增强等思路。资源包为1个PDF文件,约1.65MB,适合作为数据分析与数据研究方向的参考文献。目前已有727人学习,可帮助读者快速把握该领域的方法脉络与实验结论。
1. 基于深度学习算法的OFDM信号检测:从传统估计到神经网络的落地路径
OFDM 信号检测这件事,传统做法绕不开信道估计加均衡这一套:先发导频、估 CSI、再迫零或者 MMSE 均衡,最后判决。这套流程在 3GPP 各类信道模型下跑了很多年,成熟、可解释、复杂度可控。但一旦进了双选择性信道——高速移动带来的多普勒扩展叠加多径时延——导频开销会迅速吃掉频谱效率,MMSE 里的矩阵求逆也变成实时性瓶颈。基于深度学习算法的 OFDM 信号检测,本质上是想用数据驱动的方式,把「信道估计 + 均衡 + 判决」这条链路整体或部分替换成一个可训练的映射网络,让它在低导频密度、高多普勒、强非线性失真场景下逼近甚至超过传统接收机。
这篇笔记面向的是已经懂 OFDM 基带流程、想动手把深度学习接进接收链的工程师,以及正在做 ofdm isac 一体化波形、需要联合检测的通信方向研究者。我会按「先立住理论、再跑通最小闭环、最后讲清参数和坑」的顺序展开,代码以 PyTorch 为主,MATLAB 只作为数据生成和对照工具。读完你应该能自己搭一个从仿真数据生成到模型训练、再到误码率对比的完整流程,并且知道哪些参数一动就会翻车。
2. 为什么把检测交给神经网络:OFDM 接收链的痛点与网络选型
2.1 传统检测在双选择性信道下的三个硬伤
先说清楚为什么值得换思路。OFDM 接收端传统链路是:去 CP → FFT → 导频提取 → 信道估计(LS/MMSE/DFT 插值)→ 均衡 → 解调。这条链在慢衰落里几乎无懈可击,但双选择性信道下有三个绕不过去的问题。
第一,导频开销。块状导频为了跟踪时变信道,导频间隔要小于相干时间,梳状导频间隔要小于相干带宽。高速场景下相干时间可能只有几百微秒,导频密度上去,频谱效率直接掉两到三成。
第二,信道估计误差传播。LS 估计噪声放大,MMSE 需要信道统计先验,实际系统里这个先验往往不准,估计误差会一路传到判决,低 SNR 段误码率平台明显。
第三,MMSE 均衡的矩阵求逆。子载波数 N 大时,逐符号求逆复杂度是 O(N³),即便用近似也有可观延迟。深度学习检测器的吸引力在于:推理阶段基本是矩阵乘加,复杂度可控,而且能把信道估计、均衡、判决甚至解调揉进一个网络,端到端优化。
常见做法有两类。一类是模型驱动:把传统算法的迭代展开成网络层,比如 OAMP-Net、DetNet,每层对应一次迭代,参数可学。另一类是数据驱动:直接拿接收信号(或经过粗估计的频域信号)喂给 CNN/Transformer,输出比特或符号。前者可解释、样本效率高,后者上限高但吃数据。我一般建议从模型驱动起步,因为它在低导频场景下更稳,训练也更容易收敛。
2.2 网络结构选型:CNN、ResNet 还是 Transformer
选型要看你的输入表示。如果输入是频域接收信号排成的二维网格(子载波 × 符号),CNN 天然适合抓局部相关性,因为相邻子载波的信道响应是相关的。ResNet 在此基础上加深,缓解梯度消失,适合子载波数较大的场景(比如 N=1024 以上)。
Transformer 的优势在长程依赖,如果你的信道在频域上有跨块的相关性,或者你要做 ofdm isac 里通信与感知联合处理,注意力机制能把不同子载波、不同符号之间的关系建模进去。代价是数据量和算力需求都上一个台阶,边缘计算 深度学习 推理优化 调度 这类部署场景下要慎重。
一个务实的组合是:浅层 CNN 做特征提取,中间接几层残差块,最后全连接输出每个子载波上的比特对数似然比(LLR)。这样既保留了局部特征,又不会让参数量爆炸。下面这张表是我在几个典型配置下的选型参考。
| 场景 | 子载波数 | 推荐结构 | 参数量级 | 备注 |
|---|---|---|---|---|
| 慢衰落,导频充足 | 64~256 | 3 层 CNN | 10K~50K | 传统 MMSE 已够好,DL 收益有限 |
| 双选择性,低导频 | 256~1024 | ResNet-8 | 100K~500K | 收益最明显 |
| ISAC 联合 | 1024+ | 轻量 Transformer | 500K~2M | 需要大量仿真数据 |
| 边缘部署 | 任意 | 深度可分离卷积 | <50K | 精度换延迟 |
2.3 用 MATLAB 生成 OFDM 仿真数据集
训练检测器第一步是造数据。我一般用 MATLAB 生成,因为它的 5G Toolbox 和通信工具箱能直接给出标准信道模型,省得自己写多径和多普勒。核心是生成「接收频域信号 + 真实发送符号」的配对,网络学的是从接收信号到发送符号的映射。
% generate_ofdm_dataset.m % 生成 OFDM 双选择性信道数据集,保存为 .mat 供 PyTorch 读取 clear; rng(42); N = 256; % 子载波数 cpLen = 32; % 循环前缀长度 numSymbols = 14; % 每帧 OFDM 符号数 numFrames = 5000; % 总帧数 modOrder = 4; % QPSK snrList = -5:2:20; % SNR 扫描范围 dB % 16QAM/QPSK 星座 qpsk = (1/sqrt(2)) * [1+1j, 1-1j, -1+1j, -1-1j]; dataset = struct('rx', {}, 'tx', {}, 'snr', {}); for f = 1:numFrames snr = snrList(randi(length(snrList))); % 随机发送比特 -> QPSK 符号 bits = randi([0 1], N*numSymbols*2, 1); symIdx = bi2de(reshape(bits, 2, []).', 'left-msb') + 1; txSym = qpsk(symIdx).'; txGrid = reshape(txSym, N, numSymbols); % OFDM 调制 txTime = ifft(txGrid, N, 1); txCP = [txTime(end-cpLen+1:end, :); txTime]; txStream = txCP(:); % 双选择性信道:多径 + 多普勒 pathDelays = [0 1e-6 2e-6]; % 秒 pathGains = [0 -3 -6]; % dB maxDoppler = 300; % Hz,对应高速场景 chan = comm.RicianChannel( ... 'SampleRate', N*15e3, ... 'PathDelays', pathDelays, ... 'AveragePathGains', pathGains, ... 'MaximumDopplerShift', maxDoppler, ... 'KFactor', 0); rxStream = chan(txStream); % AWGN rxStream = awgn(rxStream, snr, 'measured'); % 接收端:去 CP + FFT rxMat = reshape(rxStream, N+cpLen, numSymbols); rxMat = rxMat(cpLen+1:end, :); rxGrid = fft(rxMat, N, 1); dataset(f).rx = single(rxGrid); dataset(f).tx = single(txGrid); dataset(f).snr = single(snr); end save('ofdm_dataset.mat', 'dataset', '-v7.3');这段脚本的逻辑:每帧随机选一个 SNR,生成 QPSK 符号,做 OFDM 调制加 CP,过 Rician 信道(K=0 即 Rayleigh),加高斯白噪声,接收端去 CP 做 FFT,得到频域接收网格。MaximumDopplerShift设 300 Hz 是模拟高速场景,PathDelays和AveragePathGains控制多径强度。保存成 v7.3 格式是因为帧数多、数据量大,老格式存不下。
参数上要注意:SampleRate必须和子载波间隔乘子载波数一致,否则多普勒和时延的物理意义就错了。numFrames建议至少 5000,不然网络在低 SNR 段学不到足够的噪声模式。SNR 范围覆盖 -5 到 20 dB,是因为检测器真正有价值的是低 SNR 段,高 SNR 段传统方法已经够好。
2.4 用 PyTorch 搭一个最小可用的检测网络
数据有了,接下来搭网络。下面这个结构是 CNN + 残差块 + LLR 输出,输入是复数频域网格,实部虚部拆成两通道。
import torch import torch.nn as nn import scipy.io as sio import numpy as np class ResidualBlock(nn.Module): """两层卷积 + 跳跃连接,缓解深层网络梯度消失""" def __init__(self, channels): super().__init__() self.conv1 = nn.Conv2d(channels, channels, 3, padding=1) self.bn1 = nn.BatchNorm2d(channels) self.conv2 = nn.Conv2d(channels, channels, 3, padding=1) self.bn2 = nn.BatchNorm2d(channels) self.relu = nn.ReLU() def forward(self, x): residual = x out = self.relu(self.bn1(self.conv1(x))) out = self.bn2(self.conv2(out)) return self.relu(out + residual) class OFDMDetector(nn.Module): """输入: (B, 2, N, T) 实部虚部双通道; 输出: (B, N*T*2) LLR""" def __init__(self, n_subcarriers=256, n_symbols=14, n_bits_per_sym=2): super().__init__() self.n_sub = n_subcarriers self.n_sym = n_symbols self.n_bits = n_bits_per_sym self.input_conv = nn.Sequential( nn.Conv2d(2, 32, 3, padding=1), nn.BatchNorm2d(32), nn.ReLU() ) self.res_blocks = nn.Sequential( ResidualBlock(32), ResidualBlock(32), ResidualBlock(32) ) # 输出每个子载波每个符号的比特 LLR self.head = nn.Conv2d(32, n_bits_per_sym, 1) def forward(self, x): # x: (B, 2, N, T) feat = self.input_conv(x) feat = self.res_blocks(feat) llr = self.head(feat) # (B, n_bits, N, T) return llr.reshape(x.size(0), -1) # 数据加载 def load_dataset(path): mat = sio.loadmat(path, simplify_cells=True) data = mat['dataset'] rx = np.stack([d['rx'] for d in data]) # (F, N, T) complex tx = np.stack([d['tx'] for d in data]) # 复数拆实部虚部 -> (F, 2, N, T) rx_real = np.stack([rx.real, rx.imag], axis=1).astype(np.float32) # QPSK 标签: 每个符号 2 bit tx_idx = np.argmin(np.abs(tx[..., None] - np.array([1+1j, 1-1j, -1+1j, -1-1j])/np.sqrt(2)), axis=-1) bits = np.stack([(tx_idx >> 1) & 1, tx_idx & 1], axis=1).astype(np.float32) return rx_real, bits if __name__ == '__main__': rx, bits = load_dataset('ofdm_dataset.mat') print('rx shape:', rx.shape, 'bits shape:', bits.shape) model = OFDMDetector() x = torch.tensor(rx[:4]) out = model(x) print('output shape:', out.shape) # 应为 (4, 256*14*2)网络逻辑:输入卷积把双通道升到 32 通道,三个残差块提取特征,1×1 卷积把通道压到每符号比特数,最后展平成 LLR 向量。残差块的作用是让网络在 256 子载波这种规模下还能加深而不退化。输出维度是N*T*n_bits,对应每个子载波每个符号的每个比特。
参数说明:n_bits_per_sym=2对应 QPSK,换 16QAM 就改成 4,同时标签生成也要改。input_conv的 32 通道是个经验值,数据量小可以降到 16,数据量大可以升到 64。残差块数量 3 是精度和延迟的折中,再往上加收益递减。
2.5 训练循环与损失函数的关键设置
检测任务本质是逐比特二分类,所以用 BCEWithLogitsLoss,不要用 MSE。MSE 在 LLR 上优化会让低 SNR 段收敛很慢。
import torch from torch.utils.data import TensorDataset, DataLoader def train_detector(rx, bits, epochs=50, batch_size=64, lr=1e-3): device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') X = torch.tensor(rx) Y = torch.tensor(bits).reshape(len(bits), -1) # (F, N*T*2) ds = TensorDataset(X, Y) dl = DataLoader(ds, batch_size=batch_size, shuffle=True) model = OFDMDetector().to(device) # L2 正则化通过 weight_decay 实现,抑制过拟合 optimizer = torch.optim.Adam(model.parameters(), lr=lr, weight_decay=1e-4) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, epochs) criterion = nn.BCEWithLogitsLoss() for ep in range(epochs): model.train() total_loss = 0 for xb, yb in dl: xb, yb = xb.to(device), yb.to(device) optimizer.zero_grad() logits = model(xb) loss = criterion(logits, yb) loss.backward() # 梯度裁剪,防止残差网络初期梯度爆炸 torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0) optimizer.step() total_loss += loss.item() scheduler.step() print(f'epoch {ep+1}, loss {total_loss/len(dl):.4f}') torch.save(model.state_dict(), 'ofdm_detector.pth') return model if __name__ == '__main__': rx, bits = load_dataset('ofdm_dataset.mat') train_detector(rx, bits)损失函数选 BCEWithLogitsLoss 是因为它把 sigmoid 和 BCE 合在一起,数值更稳。weight_decay=1e-4就是深度学习 l2 正则化 在 PyTorch 里的落地方式,防止网络记住训练集的噪声。clip_grad_norm_设 5.0 是血泪经验,残差网络在训练初期梯度容易冲高,不裁剪会直接 NaN。余弦退火让学习率平滑下降,比阶梯下降更稳。
训练完拿验证集算 BER,和 MMSE 基线对比。如果低 SNR 段 BER 比 MMSE 低一个数量级,说明网络学到了信道结构;如果只在高 SNR 段好,多半是过拟合了训练集的 SNR 分布,要检查数据生成时的 SNR 采样是否均匀。
3. 从仿真到可用:数据管线、训练策略与推理部署
3.1 数据集划分与 SNR 分层采样
上一章的数据生成是均匀采样 SNR,但实际训练时如果直接随机划分,可能出现训练集和验证集 SNR 分布不一致,导致验证 BER 虚高。我一般按 SNR 分层划分:每个 SNR 档位取 80% 训练、10% 验证、10% 测试。这样验证集能真实反映各 SNR 段的泛化。
def stratified_split(rx, bits, snr, train_ratio=0.8, val_ratio=0.1): """按 SNR 分层划分,避免分布偏移""" train_idx, val_idx, test_idx = [], [], [] for s in np.unique(snr): idx = np.where(snr == s)[0] np.random.shuffle(idx) n = len(idx) n_train = int(n * train_ratio) n_val = int(n * val_ratio) train_idx.extend(idx[:n_train]) val_idx.extend(idx[n_train:n_train+n_val]) test_idx.extend(idx[n_train+n_val:]) return (rx[train_idx], bits[train_idx], rx[val_idx], bits[val_idx], rx[test_idx], bits[test_idx])分层的好处是每个 SNR 段都有独立的测试样本,画 BER-SNR 曲线时不会因为某段样本太少而抖动。train_ratio设 0.8 是常规做法,数据量少于 2000 帧时可以降到 0.7。
3.2 迁移学习:从 QPSK 模型迁移到 16QAM
实际系统里调制阶数会变,如果每个调制方式都从头训,数据和时间成本都高。常见做法是先训 QPSK 模型,再冻结前面的卷积层,只微调输出头。因为信道特征提取部分和调制方式无关,只有最后的判决边界需要调整。
def adapt_to_16qam(qpsk_model_path, rx_16qam, bits_16qam): """加载 QPSK 预训练模型,替换输出头后微调""" model = OFDMDetector(n_bits_per_sym=2) model.load_state_dict(torch.load(qpsk_model_path)) # 替换输出头: 2 bit -> 4 bit model.head = nn.Conv2d(32, 4, 1) # 冻结前层,只训 head 和最后一个残差块 for name, param in model.named_parameters(): if 'head' not in name and 'res_blocks.2' not in name: param.requires_grad = False optimizer = torch.optim.Adam( filter(lambda p: p.requires_grad, model.parameters()), lr=5e-4) # 后续训练循环同上,epoch 可以减到 15 return model冻结策略是关键:全冻结只训 head 收敛快但精度有限,全解冻又容易过拟合小数据集。折中方案是解冻最后一个残差块加 head,学习率降到 5e-4。这样 15 个 epoch 就能达到从头训 50 个 epoch 的效果。
3.3 推理部署:ONNX 导出与延迟测量
训练完要落地,PyTorch 模型直接上生产不现实,一般导出 ONNX 再用 TensorRT 或 ONNX Runtime 推理。导出时注意输入维度要固定,动态 shape 在边缘设备上支持不好。
import torch.onnx def export_onnx(model_path, onnx_path, n_sub=256, n_sym=14): model = OFDMDetector(n_subcarriers=n_sub, n_symbols=n_sym) model.load_state_dict(torch.load(model_path)) model.eval() dummy = torch.randn(1, 2, n_sub, n_sym) torch.onnx.export( model, dummy, onnx_path, input_names=['rx_grid'], output_names=['llr'], opset_version=13, dynamic_axes=None # 固定 shape,边缘设备友好 ) print(f'exported to {onnx_path}') # 延迟测量 import onnxruntime as ort import time def measure_latency(onnx_path, n_runs=100): sess = ort.InferenceSession(onnx_path, providers=['CPUExecutionProvider']) x = np.random.randn(1, 2, 256, 14).astype(np.float32) # 预热 for _ in range(10): sess.run(None, {'rx_grid': x}) t0 = time.perf_counter() for _ in range(n_runs): sess.run(None, {'rx_grid': x}) t1 = time.perf_counter() print(f'avg latency: {(t1-t0)/n_runs*1000:.2f} ms')opset_version=13是兼容性较好的选择,再高有些边缘推理框架不支持。dynamic_axes=None固定输入 shape,能让推理引擎做更激进的图优化。延迟测量一定要先预热,第一次推理包含图加载和内存分配,不预热测出来的数会偏大好几倍。
3.4 和传统 MMSE 的公平对比方法
对比不能只看 BER 曲线,要控制变量。导频密度、信道模型、SNR 定义必须一致。我一般用同一份接收数据,分别跑 MMSE 均衡和神经网络推理,画在同一张图上。
| 对比维度 | MMSE 基线 | 深度学习检测器 |
|---|---|---|
| 导频开销 | 每帧 4 个导频符号 | 可降至 1 个或不用 |
| 信道先验 | 需要统计信息 | 从数据中学 |
| 推理复杂度 | O(N³) 求逆 | O(N·K²) 卷积 |
| 低 SNR 表现 | 误差平台明显 | 可低一个数量级 |
| 可解释性 | 强 | 弱,需额外分析 |
公平对比的前提是 MMSE 也要调到最优,包括信道估计插值方式、均衡正则项。很多人拿一个没调好的 MMSE 去比,结论不可信。另外 SNR 定义要统一用接收端每比特能量比噪声功率谱密度,不要用符号 SNR,否则两条曲线对不上。
4. 避坑与排查:训练不收敛、BER 平台、部署翻车的真实记录
4.1 损失降到 0.1 但 BER 不降
现象:训练 loss 一路降到 0.1 以下,但验证集 BER 卡在 0.1 左右不动。
原因:标签对齐错了。QPSK 的比特到符号映射有格雷码和非格雷码之分,如果 MATLAB 生成时用的映射和 PyTorch 标签生成时用的不一致,网络学的是一个错位的映射,loss 能降但 BER 下不去。
解决:统一用格雷映射,并且在数据生成脚本里显式写出星座点顺序,Python 端按同样顺序生成标签。验证方法是拿训练集的一个样本,手工算一遍 LLR 符号,看和网络输出是否同号。
4.2 低 SNR 段 BER 比 MMSE 还差
现象:高 SNR 段神经网络明显优于 MMSE,但 0 dB 以下反而更差。
原因:训练集 SNR 分布偏了。如果均匀采样 -5 到 20 dB,低 SNR 样本占比其实不低,但低 SNR 段的梯度被高 SNR 段主导,网络倾向于先学好高 SNR。另外 BCE 损失在极端噪声下梯度很小,收敛慢。
解决:对低 SNR 样本加权,或者分 SNR 段训练再融合。我一般用样本加权,权重和 SNR 成反比,让低 SNR 样本的 loss 贡献更大。另一个办法是在低 SNR 段多生成数据,比如 -5 到 5 dB 生成两倍样本。
4.3 导出 ONNX 后推理结果和 PyTorch 不一致
现象:PyTorch 里 BER 0.01,导出 ONNX 后 BER 0.05。
原因:BatchNorm 在训练和推理模式下的行为不同。导出时如果模型没设eval(),BN 会用 batch 统计量而不是滑动平均,结果自然不对。另一个常见原因是卷积的 padding 模式,PyTorch 默认零填充,某些推理框架默认可能是其他模式。
解决:导出前必须model.eval(),并且用torch.no_grad()包住导出过程。导出后用同一份输入分别跑 PyTorch 和 ONNX,逐元素比对输出,差异应该在 1e-5 以内。如果差得多,检查 opset 版本和算子支持列表。
4.4 残差块加多了反而变差
现象:从 3 个残差块加到 8 个,训练 loss 更低但验证 BER 更高。
原因:过拟合。OFDM 检测的训练数据是仿真生成的,信道模型有限,网络容量太大就会记住训练集的特定信道实现,泛化到测试集就崩。这和深度学习知识点 里讲的容量与泛化权衡是一回事。
解决:加 dropout 或者 weight_decay,或者直接减层。我的经验是 256 子载波以内 3 到 4 个残差块足够,再往上加必须同步增加数据量和信道模型多样性,否则就是负收益。
4.5 推理延迟在边缘设备上超标
现象:服务器上推理 2 ms,部署到边缘设备变成 50 ms。
原因:边缘设备的算力和内存带宽都有限,卷积层在 CPU 上可能没有优化,而且 BatchNorm 在推理时可以融合进卷积但很多框架默认不融。
解决:导出前做 BN 融合,用torch.quantization做 INT8 量化,或者换深度可分离卷积。实测 INT8 量化能把延迟降一半以上,精度损失在 0.5 dB 以内。如果还不行,就得考虑模型剪枝,把通道数从 32 降到 16。
5. 进阶技巧:用模型驱动网络把传统算法知识注入检测器
纯数据驱动的网络有个天花板:它不知道 OFDM 的物理结构,全靠数据学。模型驱动的方法把传统迭代算法展开成网络层,每一层对应一次估计或均衡迭代,层间参数可学。这样既保留了传统算法的结构先验,又能通过训练优化参数,样本效率比纯 CNN 高得多。
以 OAMP 展开为例,传统 OAMP 每次迭代做:线性估计 → 去相关 → 非线性去噪。把它展开成网络,线性估计的步长、去噪器的阈值都变成可学参数。下面是一个简化版的展开层实现。
class OAMPLayer(nn.Module): """OAMP 单层展开: 线性估计 + 可学去噪""" def __init__(self, n_sub, n_sym): super().__init__() # 可学的线性估计步长 self.step = nn.Parameter(torch.tensor(0.5)) # 可学的去噪阈值 self.threshold = nn.Parameter(torch.tensor(0.1)) self.n_sub = n_sub self.n_sym = n_sym def forward(self, y, x_est, h_est): # y: 接收信号 (B, N, T) # x_est: 当前符号估计 # h_est: 信道估计 # 线性估计: x_tmp = x_est + step * H^H (y - H x_est) residual = y - h_est * x_est x_tmp = x_est + self.step * torch.conj(h_est) * residual # 非线性去噪: 软阈值 magnitude = torch.abs(x_tmp) scale = torch.relu(magnitude - self.threshold) / (magnitude + 1e-8) x_new = x_tmp * scale return x_new class OAMPNet(nn.Module): """多层 OAMP 展开网络""" def __init__(self, n_layers=6, n_sub=256, n_sym=14): super().__init__() self.layers = nn.ModuleList([ OAMPLayer(n_sub, n_sym) for _ in range(n_layers) ]) # 初始信道估计用简单 CNN self.h_estimator = nn.Sequential( nn.Conv2d(2, 16, 3, padding=1), nn.ReLU(), nn.Conv2d(16, 2, 3, padding=1) ) def forward(self, y): # y: (B, 2, N, T) 实部虚部 h = self.h_estimator(y) h_complex = torch.complex(h[:, 0], h[:, 1]) y_complex = torch.complex(y[:, 0], y[:, 1]) x_est = torch.zeros_like(y_complex) for layer in self.layers: x_est = layer(y_complex, x_est, h_complex) # 输出 LLR return torch.cat([x_est.real, x_est.imag], dim=1)这个结构的核心思想:step和threshold不是手工调的,而是从数据里学。传统 OAMP 需要根据信道统计手动设这些参数,展开后网络自动优化。h_estimator用一个小 CNN 做粗信道估计,替代传统导频估计,这样整个接收链端到端可训。
训练这种网络要注意:层数不要太多,6 到 8 层足够,再多和纯 CNN 一样会过拟合。学习率要比纯 CNN 小,因为展开层的参数对初值敏感,我一般用 1e-4 起步。另外初始化时把step设成 0.5、threshold设成 0.1,接近传统算法的典型值,收敛更快。
验证模型驱动网络是否真的学到了东西,有个简单办法:把训练后的step和threshold打印出来,看它们是否随层数变化。如果每层参数几乎一样,说明网络没学到分层迭代的结构,退化成单层了。正常情况应该是前几层 step 大、后几层 step 小,threshold 逐层收紧,这和传统 OAMP 的收敛行为一致。
我自己的习惯是:新场景先跑纯 CNN 摸上限,再用模型驱动网络看能不能用更少数据达到接近的效果。如果模型驱动网络在 1/3 数据量下就能达到纯 CNN 的 BER,那它值得上生产,因为实际系统里标注数据永远比仿真数据贵。这套流程我在几个双选择性信道项目里反复用过,最深的教训是别一上来就堆 Transformer,先把 CNN 和展开网络跑透,大部分场景够用了。希望帮到你。
本文还有配套的精品资源,点击获取