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

摘要

MIMO(多输入多输出)技术通过多天线收发显著提升了无线通信系统的容量,是4G/5G的核心技术之一。接收端的信号检测是影响系统性能的关键环节。传统检测算法中,ZF和MMSE计算简单但性能有限,ML检测性能最优但复杂度随天线数和调制阶数指数增长,难以实际应用。近年来,深度学习被引入MIMO检测领域,本文基于深度展开网络实现了一种检测器,并与ZF、MMSE、ML进行系统对比。仿真结果表明,基于深度学习的检测器在低信噪比下与ZF MMSE相当,高信噪比下优于MMSE和ZF,检测速度远快于三种基线,验证了深度学习方法在MIMO检测中的可行性和潜力。

第1章 背景意义

1.1 研究背景

  • 移动通信技术在过去几十年经历了从1G到5G的跨越式发展,每一代技术的演进都带来了传输速率和系统容量的显著提升。MIMO技术通过在收发端配置多根天线,利用空间维度资源,在不增加频谱带宽和发射功率的前提下大幅提高系统容量和传输可靠性,已成为现代无线通信系统的核心技术之一。

  • 在MIMO系统中,信号从多根发射天线发出,经过无线信道的混合与畸变,再叠加噪声到达接收端。接收端的核心任务是从被严重污染的混合信号中恢复出各根天线发送的原始信息,这一过程称为MIMO检测。检测算法的好坏直接影响通信质量,是MIMO系统中决定性能的关键环节。

1.2 研究意义

传统的MIMO检测算法各有优劣:

  • ZF(迫零)检测:完全消除信道干扰,但会放大噪声,性能有限。

  • MMSE(最小均方误差)检测:在消除干扰和抑制噪声间寻求平衡,性能较好且计算简单,是当前实际系统中应用最广泛的方案。

  • ML(最大似然)检测:理论上性能最优,但需要穷举所有可能的发送符号组合,复杂度随天线数和调制阶数指数增长,在小规模系统中尚可接受,但面对大规模MIMO或高阶调制时变得不可行。

近年来,深度学习在图像识别、自然语言处理等领域取得了突破性进展,本文旨在实现一种基于深度展开网络的MIMO检测器,通过系统仿真与传统ZF、MMSE、ML进行全面的性能对比,验证深度学习在MIMO检测中的可行性和有效性,为后续研究提供参考。

第2章 MIMO系统模型与理论基础

2.1 MIMO信号模型

在这里插入图片描述

2.2 QPSK调制与星座映射

在这里插入图片描述

2.3 实值等效模型

在这里插入图片描述

2.4 检测问题的数学形式

在这里插入图片描述

第三章 传统检测算法

3.1 迫零(ZF)检测

在这里插入图片描述

3.2 最小均方误差(MMSE)检测

在这里插入图片描述

3.3 最大似然(ML)检测

在这里插入图片描述

第四章 基于深度学习的检测算法

4.1 深度展开的基本思想

传统迭代算法(如梯度下降法、迭代软阈值算法等)通过反复更新变量逐步逼近最优解,每一次迭代都依赖相同的数学形式,只是变量值在不断变化。深度展开的核心思想是:将迭代算法的每一次迭代视为神经网络的一层,每一层具有相同的结构,但各层的参数(如权重矩阵、偏置、步长等)可以被独立学习,而不是像传统算法那样固定不变。
在这里插入图片描述
深度展开的优势在于:

  • 结构先验:网络结构来源于成熟的迭代算法,不是随意设计的,因此训练更容易收敛。

  • 参数共享与个性化:每层的结构相同但参数不同,兼顾了通用性和灵活性。

  • 快速推理:训练完成后只需前向传播固定层数,无需迭代,推理速度快。

  • 可解释性:每一层的输出都有明确的物理意义(如当前估计值、梯度等)。

4.2 网络结构设计

在这里插入图片描述
在这里插入图片描述

4.3 训练数据

在这里插入图片描述
训练完成后,推理时只需一次前向传播即可得到输出,无需迭代或矩阵求逆。因此推理速度非常快,尤其适合GPU并行处理大量数据。

第五章 仿真结果与分析

5.1 仿真参数

MIMO参数名称 参数值
发射天线数 N t N_t Nt 4
接收天线数 N r N_r Nr 4
调制方式 QPSK(格雷映射)
信道模型 平坦瑞利衰落,每帧独立生成
信道矩阵分布 C N ( 0 , 1 ) \mathcal{CN}(0,1) CN(0,1),每元素独立同分布
信道估计 理想(接收端已知精确 H \mathbf{H} H
噪声模型 复加性高斯白噪声(AWGN)
每SNR点测试帧数 2000
参数名称 参数值
训练样本数 100,000
验证样本数 10,000
批量大小(batch size) 512
训练轮次(epochs) 40
网络层数 L L L 10
隐层维度 D h D_h Dh 128
状态维度 D s D_s Ds 64
优化器 Adam
学习率 1 × 10 − 3 1 \times 10^{-3} 1×103

5.2 仿真图分析

在这里插入图片描述
在这里插入图片描述

可以看到:

  • 基于深度学习的检测方法性能优于ZF和mmse,但是差于ML;
  • 推理速度显著快于其他三种算法。

部分代码:



from __future__ import annotations

from pathlib import Path
import time

import pandas as pd
import torch
from torch import nn
from torch.utils.data import DataLoader
from tqdm.auto import tqdm

from .models import DetNet
from .modulation import (
    bits_to_symbols,
    nearest_constellation_points,
    real_vector_to_complex,
    symbols_to_bits,
    get_constellation,
)
from .utils import ensure_dir


def symbol_penalty(output: torch.Tensor, modulation: str) -> torch.Tensor:

    B, D = output.shape
    nt = D // 2
    real = output[:, :nt]
    imag = output[:, nt:]
    symbols_complex = real + 1j * imag  # (B, nt)
    constellation, _ = get_constellation(modulation)
    const = torch.tensor(constellation, dtype=output.dtype, device=output.device)  # (M,)
    # Distance squared (B, nt, M)
    diff = symbols_complex.unsqueeze(-1) - const
    dist_sq = torch.abs(diff) ** 2
    min_dist_sq, _ = torch.min(dist_sq, dim=-1)  # (B, nt)
    return min_dist_sq.mean()


def detnet_loss(
    outputs: torch.Tensor | list[torch.Tensor],
    target: torch.Tensor,
    modulation: str,
    penalty_weight: float = 0.01,
) -> torch.Tensor:
    """Weighted MSE loss + symbol distance penalty."""
    criterion = nn.MSELoss()
    if isinstance(outputs, list):
        weights = torch.arange(1, len(outputs) + 1, dtype=target.dtype, device=target.device)
        weights = weights / weights.sum()
        mse_loss = torch.zeros((), dtype=target.dtype, device=target.device)
        for weight, output in zip(weights, outputs):
            mse_loss = mse_loss + weight * criterion(output, target)
        final_output = outputs[-1]
    else:
        mse_loss = criterion(outputs, target)
        final_output = outputs

    penalty = symbol_penalty(final_output, modulation)
    return mse_loss + penalty_weight * penalty


def train_one_epoch(
    model: DetNet,
    loader: DataLoader,
    optimizer: torch.optim.Optimizer,
    device: torch.device,
    modulation: str,
    grad_clip: float = 5.0,
    penalty_weight: float = 0.01,
) -> float:
    model.train()
    total_loss = 0.0
    total_items = 0
    for batch in tqdm(loader, desc="train", leave=False):
        h_real = batch["H_real"].to(device)
        y_real = batch["y_real"].to(device)
        x_real = batch["x_real"].to(device)
        noise_var = batch["noise_var"].to(device=device, dtype=x_real.dtype)

        optimizer.zero_grad(set_to_none=True)
        outputs = model(h_real, y_real, noise_var=noise_var, return_all=True)
        loss = detnet_loss(outputs, x_real, modulation, penalty_weight=penalty_weight)
        loss.backward()
        if grad_clip > 0.0:
            torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip)
        optimizer.step()
        batch_size = int(x_real.shape[0])
        total_loss += float(loss.item()) * batch_size
        total_items += batch_size
    return total_loss / max(total_items, 1)


def _symbol_error_counts_from_output(
    output: torch.Tensor,
    bits: torch.Tensor,
    modulation: str,
) -> tuple[int, int, int, int]:
    estimates = output.detach().cpu().numpy()
    true_bits_batch = bits.detach().cpu().numpy()
    bit_errors = total_bits = symbol_errors = total_symbols = 0
    for estimate_real, true_bits in zip(estimates, true_bits_batch):
        x_hat = nearest_constellation_points(real_vector_to_complex(estimate_real), modulation)
        x_true = bits_to_symbols(true_bits.reshape(-1), modulation)
        estimated_bits = symbols_to_bits(x_hat, modulation)
        true_bits_flat = true_bits.reshape(-1)
        bit_errors += int((estimated_bits != true_bits_flat).sum())
        total_bits += int(true_bits_flat.size)
        symbol_errors += int((abs(x_hat - x_true) > 1e-8).sum())
        total_symbols += int(x_true.size)
    return bit_errors, total_bits, symbol_errors, total_symbols


@torch.no_grad()
def validate(
    model: DetNet,
    loader: DataLoader,
    device: torch.device,
    modulation: str,
    max_metric_samples: int | None = None,
) -> dict[str, float]:
    model.eval()
    total_loss = 0.0
    total_items = 0
    bit_errors = total_bits = symbol_errors = total_symbols = 0
    metric_items = 0
    for batch in tqdm(loader, desc="validate", leave=False):
        h_real = batch["H_real"].to(device)
        y_real = batch["y_real"].to(device)
        x_real = batch["x_real"].to(device)
        noise_var = batch["noise_var"].to(device=device, dtype=x_real.dtype)

        output = model(h_real, y_real, noise_var=noise_var, return_all=False)
        # Validation loss uses MSE only (no penalty)
        loss = nn.MSELoss()(output, x_real)  # 简单MSE,不加权重
        batch_size = int(x_real.shape[0])
        total_loss += float(loss.item()) * batch_size
        total_items += batch_size

        remaining = None if max_metric_samples is None else max_metric_samples - metric_items
        if remaining is None or remaining > 0:
            metric_count = batch_size if remaining is None else min(batch_size, remaining)
            b_err, b_total, s_err, s_total = _symbol_error_counts_from_output(
                output[:metric_count],
                batch["bits"][:metric_count],
                modulation,
            )
            bit_errors += b_err
            total_bits += b_total
            symbol_errors += s_err
            total_symbols += s_total
            metric_items += metric_count

    return {
        "val_loss": total_loss / max(total_items, 1),
        "val_ber": bit_errors / max(total_bits, 1),
        "val_ser": symbol_errors / max(total_symbols, 1),
        "val_metric_samples": float(metric_items),
    }


def save_checkpoint(
    model: DetNet,
    checkpoint_path: str | Path,
    epoch: int,
    val_loss: float,
    metadata: dict,
) -> None:
    path = Path(checkpoint_path)
    ensure_dir(path.parent)
    payload = {
        "model_state_dict": model.state_dict(),
        "epoch": int(epoch),
        "val_loss": float(val_loss),
        "metadata": metadata,
    }
    torch.save(payload, path)


def train_detnet_model(
    model: DetNet,
    train_loader: DataLoader,
    val_loader: DataLoader,
    epochs: int,
    learning_rate: float,
    device: torch.device,
    checkpoint_path: str | Path,
    metadata: dict,
    modulation: str,
    val_ber_samples: int | None = 2000,
    penalty_weight: float = 0.01,
) -> pd.DataFrame:
    optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate)
    model.to(device)
    rows: list[dict] = []
    best_val_loss = float("inf")
    started_at = time.perf_counter()
    for epoch in range(1, epochs + 1):
        train_loss = train_one_epoch(
            model, train_loader, optimizer, device,
            modulation=modulation,
            penalty_weight=penalty_weight,
        )
        metrics = validate(model, val_loader, device, modulation=modulation, max_metric_samples=val_ber_samples)
        val_loss = metrics["val_loss"]
        elapsed_seconds = time.perf_counter() - started_at
        rows.append(
            {
                "epoch": epoch,
                "train_loss": train_loss,
                "val_loss": val_loss,
                "val_ber": metrics["val_ber"],
                "val_ser": metrics["val_ser"],
                "val_metric_samples": int(metrics["val_metric_samples"]),
                "learning_rate": learning_rate,
                "elapsed_seconds": elapsed_seconds,
                "device": str(device),
            }
        )
        if val_loss < best_val_loss:
            best_val_loss = val_loss
            save_checkpoint(
                model=model,
                checkpoint_path=checkpoint_path,
                epoch=epoch,
                val_loss=val_loss,
                metadata=metadata,
            )
        print(
            f"epoch={epoch} train_loss={train_loss:.6f} val_loss={val_loss:.6f} "
            f"val_ber={metrics['val_ber']:.6g} val_ser={metrics['val_ser']:.6g}"
        )
    return pd.DataFrame(rows)

5.3 总结

本文系统性地研究了基于深度展开网络的 MIMO 检测算法,并与传统 ZF、MMSE、ML 检测算法进行了对比,基于深度学习的 MIMO 检测方法在性能与复杂度之间取得了良好的平衡,为未来大规模 MIMO 系统的高效检测提供了一条有前景的技术路径。后续工作可探索更复杂的信道模型、更高阶的调制方式,以及网络结构的进一步优化。

仿真代码可见文末VX公众号(包含往期博客所有代码),所见即所得

更多推荐