Files
Graduation-Design/Graduation Design/Massive_Mimo_Beamforming_Testing_Platform.py

626 lines
24 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import math
import time
import numpy as np
import tkinter as tk
from tkinter import ttk, messagebox, filedialog
import torch
import matplotlib as mpl
from matplotlib.figure import Figure
from matplotlib.backends.backend_tkagg import FigureCanvasTkAgg, NavigationToolbar2Tk
from matplotlib import font_manager
# =========================================================
# 0. 全局设置
# =========================================================
def configure_matplotlib_fonts():
"""自动选择可用中文字体,修复中文显示问题。"""
candidate_fonts = [
"Microsoft YaHei",
"SimHei",
"Noto Sans CJK SC",
"Noto Serif CJK SC",
"WenQuanYi Zen Hei",
"AR PL UMing CN",
"PingFang SC",
"Heiti SC",
"STHeiti",
"Songti SC",
]
installed = {f.name for f in font_manager.fontManager.ttflist}
selected = None
for name in candidate_fonts:
if name in installed:
selected = name
break
mpl.rcParams["axes.unicode_minus"] = False
mpl.rcParams["mathtext.fontset"] = "stix"
if selected:
mpl.rcParams["font.sans-serif"] = [selected]
else:
mpl.rcParams["font.sans-serif"] = ["DejaVu Sans"]
configure_matplotlib_fonts()
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
torch.manual_seed(2024)
np.random.seed(2024)
L_SNAPSHOTS_DEFAULT = 100
MC_DEFAULT = 100
# 颜色表:修复 METHOD_COLORS 未定义
METHOD_COLORS = {
"Analog": "#77AC30",
"Digital": "#000000",
"DFT": "#4DBEEE",
"OMP": "#D95319",
"TSC-OMP": "#0072BD",
}
# =========================================================
# 1. 算法引擎
# =========================================================
class BeamEngine:
def __init__(self):
self.device = device
self.N_ant = 24
self.N_RF = 4
self.noise_power = 1.0
self.INR_linear = 10 ** (30 / 10)
# ---------- 阵列模型 ----------
def get_positions_lambda(self, array_type: str, n_ant: int):
array_type = array_type.upper()
if array_type == "ULA":
pos = 0.5 * np.arange(-(n_ant - 1) / 2, (n_ant - 1) / 2 + 1)
return pos.astype(np.float32)
# 24 阵元 NULA:沿用现有主线代码
if array_type == "NULA" and n_ant == 24:
fc = 4.9e9
lam = 3e8 / fc
a = np.arange(0, 11 * 43 + 1, 43)
b = np.array([11 * 43 + 63])
c = 11 * 43 + 63 + np.arange(43, 11 * 43 + 1, 43)
delta_d = 1e-3 * np.concatenate((a, b, c))
pos = delta_d / lam
return pos.astype(np.float32)
# 兜底广义 NULA
base = 0.5 * np.arange(n_ant)
jitter = 0.08 * np.sin(np.linspace(0, 3 * np.pi, n_ant))
pos = base + jitter
pos = pos - np.mean(pos)
return pos.astype(np.float32)
def gen_a(self, theta_deg, array_type="ULA"):
pos = self.get_positions_lambda(array_type, self.N_ant)
pos_t = torch.tensor(pos, dtype=torch.float32, device=self.device).unsqueeze(1)
theta_rad = torch.deg2rad(torch.tensor(theta_deg, dtype=torch.float32, device=self.device))
if array_type.upper() == "ULA":
return torch.exp(1j * 2 * torch.pi * pos_t * torch.sin(theta_rad))
return torch.exp(-1j * 2 * torch.pi * pos_t * torch.sin(theta_rad))
def build_dictionary(self, angle_grid, array_type="ULA"):
return torch.cat(
[self.gen_a(float(th), array_type) / math.sqrt(self.N_ant) for th in angle_grid],
dim=1
)
# ---------- 子程序 ----------
def extract_omp(self, w_opt, a_dic):
res = w_opt.clone()
f_rf = None
for _ in range(self.N_RF):
proj = a_dic.mH @ res
idx = torch.argmax(torch.abs(proj))
atom = a_dic[:, idx:idx + 1]
f_rf = atom if f_rf is None else torch.cat((f_rf, atom), dim=1)
res = w_opt - f_rf @ (torch.linalg.pinv(f_rf) @ w_opt)
return f_rf
def make_prior_broad(self, theta_prior, window_half_width, array_type="ULA"):
a_prior_broad = torch.zeros((self.N_ant, 1), dtype=torch.complex64, device=self.device)
for tb in torch.arange(theta_prior - window_half_width,
theta_prior + window_half_width + 0.1,
0.5,
device=self.device):
a_prior_broad += self.gen_a(float(tb), array_type)
a_prior_broad /= torch.norm(a_prior_broad)
return a_prior_broad
def calc_hybrid_weight(self, f_rf, a_target, r_in):
a_eff = f_rf.mH @ a_target
r_eff = f_rf.mH @ r_in @ f_rf
r_eff = r_eff + (0.05 * torch.trace(r_eff).real / self.N_RF) * torch.eye(self.N_RF, device=self.device)
inv_r_eff = torch.linalg.inv(r_eff)
f_bb = (inv_r_eff @ a_eff) / (a_eff.mH @ inv_r_eff @ a_eff)
w = f_rf @ f_bb
return w / torch.norm(w)
def build_all_structures(self, target_angle, clutter_angle, array_type="ULA", window_half_width=5.0):
theta_prior = round(target_angle)
a_target = self.gen_a(target_angle, array_type)
a_clutter = self.gen_a(clutter_angle, array_type)
# 干扰+噪声先验协方差
r_interference_prior = (
self.noise_power * torch.eye(self.N_ant, device=self.device)
+ self.INR_linear * (a_clutter @ a_clutter.mH)
)
a_prior_broad = self.make_prior_broad(theta_prior, window_half_width, array_type)
w_digital_prior = torch.linalg.inv(r_interference_prior) @ a_prior_broad
w_digital_prior /= torch.norm(w_digital_prior)
# 传统 OMP / TSC-OMP 字典
dic_angles_full = np.arange(-90, 90.5, 0.5)
dic_angles_tsc = np.arange(
theta_prior - window_half_width,
theta_prior + window_half_width + 0.1,
0.1
)
a_dic_full = self.build_dictionary(dic_angles_full, array_type)
a_dic_tsc = self.build_dictionary(dic_angles_tsc, array_type)
# DFT 射频矩阵
angles_dft = torch.linspace(
theta_prior - window_half_width,
theta_prior + window_half_width,
self.N_RF,
device=self.device
)
f_rf_dft = torch.cat(
[self.gen_a(float(th), array_type) / math.sqrt(self.N_ant) for th in angles_dft],
dim=1
)
f_rf_omp = self.extract_omp(w_digital_prior, a_dic_full)
f_rf_tsc = self.extract_omp(w_digital_prior, a_dic_tsc)
# 静态理想协方差,用于方向图和纯数字权值
sig_power_static = 10 ** (0 / 10)
r_in_static = (
self.noise_power * torch.eye(self.N_ant, device=self.device)
+ sig_power_static * (a_target @ a_target.mH)
+ self.INR_linear * (a_clutter @ a_clutter.mH)
)
r_in_static_dl = r_in_static + (
0.01 * torch.trace(r_in_static).real / self.N_ant
) * torch.eye(self.N_ant, device=self.device)
w_analog = a_target / math.sqrt(self.N_ant)
inv_r_stat = torch.linalg.inv(r_in_static_dl)
w_digital = (inv_r_stat @ a_target) / (a_target.mH @ inv_r_stat @ a_target)
w_digital /= torch.norm(w_digital)
# DFT / OMP / TSC-OMP 混合
w_dft = self.calc_hybrid_weight(f_rf_dft, a_target, r_in_static)
w_omp = self.calc_hybrid_weight(f_rf_omp, a_target, r_in_static)
w_tsc = self.calc_hybrid_weight(f_rf_tsc, a_target, r_in_static)
return {
"theta_prior": theta_prior,
"a_target": a_target,
"a_clutter": a_clutter,
"w_analog": w_analog,
"w_digital": w_digital,
"w_dft": w_dft,
"w_omp": w_omp,
"w_tsc": w_tsc,
"f_rf_dft": f_rf_dft,
"f_rf_omp": f_rf_omp,
"f_rf_tsc": f_rf_tsc,
}
# ---------- 主功能1:波束方向图 ----------
def beam_patterns(self, target_angle, clutter_angle, array_type="ULA",
methods=None, window_half_width=5.0):
if methods is None:
methods = ["Analog", "Digital", "DFT", "OMP", "TSC-OMP"]
structures = self.build_all_structures(target_angle, clutter_angle, array_type, window_half_width)
theta_scan = torch.arange(-90, 90.1, 0.1, device=self.device)
a_scan = torch.cat([self.gen_a(float(th), array_type) for th in theta_scan], dim=1)
weights = {
"Analog": structures["w_analog"],
"Digital": structures["w_digital"],
"DFT": structures["w_dft"],
"OMP": structures["w_omp"],
"TSC-OMP": structures["w_tsc"],
}
result = {"theta_scan": theta_scan.cpu().numpy()}
for method in methods:
w = weights[method]
gain = torch.abs(a_scan.mH @ w) ** 2
bp = 10 * torch.log10((gain / torch.max(gain)).squeeze())
result[method] = bp.cpu().numpy()
return result
# ---------- 主功能2RMSE 曲线 ----------
def rmse_curves(self, target_angle, clutter_angle, array_type="ULA",
methods=None, window_half_width=5.0,
mc=100, snapshots=100, snr_values=None):
"""
RMSE 计算逻辑参考:
- ULA: RMSE_Test_ula(2).py
- NULA: RMSE_Test_nula(2).py
其中 DFT / OMP / TSC-OMP 均采用先验宽波束 + 低维 Capon 超分辨。
"""
if methods is None:
methods = ["Analog", "Digital", "DFT", "OMP", "TSC-OMP"]
if snr_values is None:
snr_values = np.arange(-10, 22, 2)
structures = self.build_all_structures(target_angle, clutter_angle, array_type, window_half_width)
theta_prior = structures["theta_prior"]
a_target = structures["a_target"]
a_clutter = structures["a_clutter"]
scan_grid = torch.arange(
theta_prior - window_half_width,
theta_prior + window_half_width + 0.01,
0.01,
device=self.device
)
a_scan = torch.cat([self.gen_a(float(th), array_type) for th in scan_grid], dim=1)
curves = {m: [] for m in methods}
a_eff_dft = structures["f_rf_dft"].mH @ a_scan
a_eff_omp = structures["f_rf_omp"].mH @ a_scan
a_eff_tsc = structures["f_rf_tsc"].mH @ a_scan
for snr in snr_values:
sig_power = 10 ** (snr / 10)
s_t = np.sqrt(sig_power / 2) * (
torch.randn(mc, 1, snapshots, device=self.device) +
1j * torch.randn(mc, 1, snapshots, device=self.device)
)
s_c = np.sqrt(self.INR_linear / 2) * (
torch.randn(mc, 1, snapshots, device=self.device) +
1j * torch.randn(mc, 1, snapshots, device=self.device)
)
n_t = np.sqrt(self.noise_power / 2) * (
torch.randn(mc, self.N_ant, snapshots, device=self.device) +
1j * torch.randn(mc, self.N_ant, snapshots, device=self.device)
)
x = a_target.unsqueeze(0) * s_t + a_clutter.unsqueeze(0) * s_c + n_t
r_in = torch.bmm(x, x.mH) / snapshots
# 纯模拟:Bartlett
if "Analog" in methods:
p_ana = torch.einsum('si, bij, js -> bs', a_scan.conj().T, r_in, a_scan).abs()
ang_ana = scan_grid[torch.argmax(p_ana, dim=1)]
curves["Analog"].append(torch.sqrt(torch.mean((ang_ana - target_angle) ** 2)).item())
# 纯数字:全维 Capon / MVDR
if "Digital" in methods:
trace_r_in = torch.diagonal(r_in, dim1=-2, dim2=-1).sum(-1).real
r_in_dl = r_in + (
0.05 * trace_r_in / self.N_ant
).view(-1, 1, 1) * torch.eye(self.N_ant, device=self.device).unsqueeze(0)
inv_r_in_dl = torch.linalg.inv(r_in_dl)
p_dig = 1.0 / torch.abs(torch.einsum('si, bij, js -> bs', a_scan.conj().T, inv_r_in_dl, a_scan))
ang_dig = scan_grid[torch.argmax(p_dig, dim=1)]
curves["Digital"].append(torch.sqrt(torch.mean((ang_dig - target_angle) ** 2)).item())
def process_hybrid(f_rf, a_eff):
y = f_rf.mH.unsqueeze(0) @ x
r = (y @ y.mH) / snapshots
r = r + (
0.05 * torch.diagonal(r, dim1=-2, dim2=-1).sum(-1).view(-1, 1, 1) / self.N_RF
) * torch.eye(self.N_RF, device=self.device).unsqueeze(0)
inv_r = torch.linalg.inv(r)
p = 1.0 / torch.abs(torch.einsum('si, bij, js -> bs', a_eff.conj().T, inv_r, a_eff))
ang = scan_grid[torch.argmax(p, dim=1)]
return torch.sqrt(torch.mean((ang - target_angle) ** 2)).item()
if "DFT" in methods:
curves["DFT"].append(process_hybrid(structures["f_rf_dft"], a_eff_dft))
if "OMP" in methods:
curves["OMP"].append(process_hybrid(structures["f_rf_omp"], a_eff_omp))
if "TSC-OMP" in methods:
curves["TSC-OMP"].append(process_hybrid(structures["f_rf_tsc"], a_eff_tsc))
curves["snr"] = np.array(snr_values)
return curves
# =========================================================
# 2. 绘图面板
# =========================================================
class PlotPanel(ttk.Frame):
def __init__(self, master):
super().__init__(master)
self.figure = Figure(figsize=(8.6, 5.8), dpi=100)
self.canvas = FigureCanvasTkAgg(self.figure, master=self)
self.canvas_widget = self.canvas.get_tk_widget()
self.canvas_widget.pack(fill="both", expand=True)
self.toolbar = NavigationToolbar2Tk(self.canvas, self, pack_toolbar=False)
self.toolbar.update()
self.toolbar.pack(fill="x")
def redraw(self):
self.canvas.draw_idle()
def save_figure(self):
path = filedialog.asksaveasfilename(
defaultextension=".png",
filetypes=[("PNG 图片", "*.png"), ("PDF 文件", "*.pdf"), ("SVG 文件", "*.svg")]
)
if path:
self.figure.savefig(path, dpi=300, bbox_inches="tight")
return path
return None
# =========================================================
# 3. GUI 主体
# =========================================================
class HBFWorkbench(tk.Tk):
def __init__(self):
super().__init__()
self.title("大规模阵列 MIMO 波束优化综合平台")
self.geometry("1450x900")
self.minsize(1200, 760)
self.engine = BeamEngine()
self._setup_style()
self.status_var = tk.StringVar(value=f"就绪 | 设备: {device}")
self._build_ui()
def _setup_style(self):
style = ttk.Style(self)
try:
style.theme_use("clam")
except Exception:
pass
style.configure("TNotebook.Tab", padding=(14, 8), font=("Microsoft YaHei", 10))
style.configure("Header.TLabel", font=("Microsoft YaHei", 12, "bold"))
style.configure("Accent.TButton", font=("Microsoft YaHei", 10, "bold"))
def _build_ui(self):
main = ttk.Frame(self)
main.pack(fill="both", expand=True, padx=10, pady=10)
left = ttk.Frame(main)
left.pack(side="left", fill="y", padx=(0, 8))
right = ttk.Frame(main)
right.pack(side="right", fill="both", expand=True)
# ---------- 参数卡 ----------
card1 = ttk.LabelFrame(left, text="场景参数")
card1.pack(fill="x", pady=(0, 8))
self.var_target = tk.DoubleVar(value=5.0)
self.var_clutter = tk.DoubleVar(value=-15.0)
self.var_window = tk.DoubleVar(value=5.0)
self.var_mc = tk.IntVar(value=MC_DEFAULT)
self.var_snap = tk.IntVar(value=L_SNAPSHOTS_DEFAULT)
self.var_array = tk.StringVar(value="ULA")
self.var_function = tk.StringVar(value="beam")
self._add_labeled_entry(card1, "目标角度 (°)", self.var_target)
self._add_labeled_entry(card1, "干扰角度 (°)", self.var_clutter)
self._add_labeled_entry(card1, "波门半宽 (°)", self.var_window)
self._add_labeled_entry(card1, "Monte Carlo", self.var_mc)
self._add_labeled_entry(card1, "快拍数", self.var_snap)
row = ttk.Frame(card1)
row.pack(fill="x", pady=4)
ttk.Label(row, text="阵列类型", width=12).pack(side="left")
ttk.Combobox(row, textvariable=self.var_array, values=["ULA", "NULA"], state="readonly", width=12).pack(side="left")
# ---------- 功能卡 ----------
card2 = ttk.LabelFrame(left, text="功能选择")
card2.pack(fill="x", pady=(0, 8))
ttk.Radiobutton(card2, text="波束方向图显示", variable=self.var_function, value="beam").pack(anchor="w")
ttk.Radiobutton(card2, text="测角 RMSE 曲线", variable=self.var_function, value="rmse").pack(anchor="w")
# ---------- 算法卡 ----------
card3 = ttk.LabelFrame(left, text="算法选择(可多选)")
card3.pack(fill="x", pady=(0, 8))
self.alg_vars = {
"Analog": tk.BooleanVar(value=True),
"Digital": tk.BooleanVar(value=True),
"DFT": tk.BooleanVar(value=True),
"OMP": tk.BooleanVar(value=True),
"TSC-OMP": tk.BooleanVar(value=True),
}
labels = {
"Analog": "纯模拟",
"Digital": "纯数字(MVDR)",
"DFT": "DFT",
"OMP": "传统OMP",
"TSC-OMP": "TSC-OMP",
}
for k in ["Analog", "Digital", "DFT", "OMP", "TSC-OMP"]:
ttk.Checkbutton(card3, text=labels[k], variable=self.alg_vars[k]).pack(anchor="w")
# ---------- 按钮 ----------
btns = ttk.Frame(left)
btns.pack(fill="x", pady=(4, 8))
ttk.Button(btns, text="运行", style="Accent.TButton", command=self.run_main_task).pack(fill="x", pady=3)
ttk.Button(btns, text="保存当前图", command=lambda: self._save_panel(self.main_plot)).pack(fill="x", pady=3)
# ---------- 说明 ----------
note = ttk.LabelFrame(left, text="说明")
note.pack(fill="x", pady=(0, 8))
ttk.Label(
note,
text="• 设置好您所需的参数,算法可以多选\n"
"• 选择您所需要的图表,有波束方向图和测角RMSE比较图可选\n"
"• 单击“运行”即可输出所需图表,可以随时保存",
justify="left"
).pack(anchor="w", padx=4, pady=4)
# ---------- 绘图区域 ----------
self.main_plot = PlotPanel(right)
self.main_plot.pack(fill="both", expand=True)
# ---------- 状态栏 ----------
status = ttk.Label(self, textvariable=self.status_var, anchor="w")
status.pack(fill="x", padx=10, pady=(0, 8))
def _add_labeled_entry(self, master, label, var, width=12):
row = ttk.Frame(master)
row.pack(fill="x", pady=4)
ttk.Label(row, text=label, width=12).pack(side="left")
ttk.Entry(row, textvariable=var, width=width).pack(side="left")
def _save_panel(self, panel):
path = panel.save_figure()
if path:
self.status_var.set(f"已保存图像: {path}")
def get_selected_methods(self):
return [k for k, v in self.alg_vars.items() if v.get()]
def _set_busy(self, msg):
self.config(cursor="watch")
self.status_var.set(msg)
self.update_idletasks()
def _set_idle(self, msg="就绪"):
self.config(cursor="")
self.status_var.set(msg)
self.update_idletasks()
def run_main_task(self):
try:
target = float(self.var_target.get())
clutter = float(self.var_clutter.get())
array_type = self.var_array.get()
func = self.var_function.get()
window = float(self.var_window.get())
mc = int(self.var_mc.get())
snapshots = int(self.var_snap.get())
methods = self.get_selected_methods()
if not methods:
messagebox.showwarning("提示", "请至少选择一种算法。")
return
t0 = time.time()
self._set_busy("正在运行,请稍候...")
fig = self.main_plot.figure
fig.clear()
ax = fig.add_subplot(111)
if func == "beam":
data = self.engine.beam_patterns(target, clutter, array_type, methods, window)
theta = data["theta_scan"]
name_map = {
"Analog": "纯模拟",
"Digital": "纯数字(MVDR)",
"DFT": "DFT",
"OMP": "传统OMP",
"TSC-OMP": "TSC-OMP",
}
ls_map = {
"Analog": "-.",
"Digital": "-",
"DFT": "--",
"OMP": ":",
"TSC-OMP": "-",
}
for m in methods:
ax.plot(
theta,
data[m],
linestyle=ls_map[m],
color=METHOD_COLORS[m],
linewidth=2.4,
label=name_map[m]
)
ax.axvline(target, color="g", linestyle="--", linewidth=1.8, label=f"目标 ({target:.2f}°)")
ax.axvline(clutter, color="r", linestyle="--", linewidth=1.8, label=f"干扰 ({clutter:.2f}°)")
ax.set_xlim([-60, 60])
ax.set_ylim([-60, 2])
ax.set_title(f"波束方向图 ({array_type})", fontsize=13, fontweight="bold")
ax.set_xlabel("角度 θ (°)")
ax.set_ylabel("归一化增益 (dB)")
ax.grid(True, ls="--", alpha=0.6)
ax.legend(loc="lower center", ncol=min(4, len(methods) + 2), fontsize=9)
else:
curves = self.engine.rmse_curves(
target_angle=target,
clutter_angle=clutter,
array_type=array_type,
methods=methods,
window_half_width=window,
mc=mc,
snapshots=snapshots,
)
name_map = {
"Analog": "纯模拟",
"Digital": "纯数字(MVDR)",
"DFT": "DFT",
"OMP": "传统OMP",
"TSC-OMP": "TSC-OMP",
}
marker_map = {
"Analog": "v",
"Digital": "o",
"DFT": "d",
"OMP": "*",
"TSC-OMP": "s",
}
for m in methods:
ax.semilogy(
curves["snr"],
curves[m],
marker=marker_map[m],
color=METHOD_COLORS[m],
linewidth=2.2,
label=name_map[m]
)
ax.set_title(f"测角 RMSE-SNR 曲线 ({array_type})", fontsize=13, fontweight="bold")
ax.set_xlabel("输入 SNR (dB)")
ax.set_ylabel("RMSE (°)")
ax.set_ylim([1e-3, 20])
ax.grid(True, which="both", ls="--", alpha=0.6)
ax.legend()
fig.tight_layout()
self.main_plot.redraw()
self._set_idle(f"运行完成,用时 {time.time() - t0:.2f}s | 设备: {device}")
except Exception as e:
self._set_idle("运行失败")
messagebox.showerror("错误", str(e))
if __name__ == "__main__":
app = HBFWorkbench()
app.mainloop()