上传文件至「Graduation Design」
This commit is contained in:
@@ -0,0 +1,153 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.optim as optim
|
||||
import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
import time
|
||||
|
||||
# ==========================================
|
||||
# 0. 全局设置 (对标你的 -35° 强杂波挑战)
|
||||
# ==========================================
|
||||
plt.rcParams['font.sans-serif'] = ['Microsoft YaHei']
|
||||
plt.rcParams['axes.unicode_minus'] = False
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
print(f"========== 启动 TSC-DL 扇区聚焦引擎 ==========\n使用设备: {device}")
|
||||
|
||||
N_ant, N_RF, d_lambda = 24, 4, 0.5
|
||||
noise_power = 1.0
|
||||
INR_linear = 10**(30/10) # 30dB 强杂波
|
||||
|
||||
#核心测试场景
|
||||
test_t_ang, test_c_ang = 15.0, -25.0
|
||||
|
||||
def gen_a_batch(theta_deg_tensor):
|
||||
theta_rad = torch.deg2rad(theta_deg_tensor)
|
||||
n_idx = torch.arange(-(N_ant-1)/2, (N_ant-1)/2 + 0.1, device=device).unsqueeze(0)
|
||||
phases = 2 * torch.pi * d_lambda * n_idx * torch.sin(theta_rad.unsqueeze(1))
|
||||
return torch.exp(1j * phases).unsqueeze(2)
|
||||
|
||||
# ==========================================
|
||||
# 1. 复数域正交网络 (彻底消灭相位梯度死锁)
|
||||
# ==========================================
|
||||
class SectorHBFNet(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.net = nn.Sequential(
|
||||
nn.Linear(48, 512), nn.GELU(),
|
||||
nn.Linear(512, 512), nn.GELU(),
|
||||
# 【核心架构升级】:不直接输出相位!输出复数的实部和虚部 (2个通道)
|
||||
nn.Linear(512, N_ant * N_RF * 2)
|
||||
)
|
||||
|
||||
def forward(self, w_opt):
|
||||
w_flat = w_opt.squeeze(2)
|
||||
x = torch.cat([w_flat.real, w_flat.imag], dim=1)
|
||||
out = self.net(x).view(-1, N_ant, N_RF, 2)
|
||||
|
||||
complex_out = out[..., 0] + 1j * out[..., 1]
|
||||
# 【物理强制约束】:将复数除以自身的模长,绝对精准地满足恒模约束,且梯度无比丝滑!
|
||||
F_RF = complex_out / (torch.abs(complex_out) + 1e-8) / np.sqrt(N_ant)
|
||||
return F_RF
|
||||
|
||||
# ==========================================
|
||||
# 2. TSC-DL 扇区聚焦训练
|
||||
# ==========================================
|
||||
model = SectorHBFNet().to(device)
|
||||
optimizer = optim.Adam(model.parameters(), lr=0.002)
|
||||
|
||||
Batch_Size, Epochs = 500, 1500 #训练次数
|
||||
loss_history = []
|
||||
|
||||
print("启动 TSC-DL 训练")
|
||||
start_time = time.time()
|
||||
for epoch in range(Epochs):
|
||||
model.train()
|
||||
optimizer.zero_grad()
|
||||
|
||||
# 【与 TSC-OMP 呼应的核心逻辑】:扇区聚焦训练!
|
||||
# 不再全空域瞎猜,而是在目标和杂波的局部战术扇区(±3°)内生成海量数据进行特训
|
||||
theta_t = test_t_ang + (torch.rand(Batch_Size, device=device) - 0.5) * 6.0
|
||||
theta_c = test_c_ang + (torch.rand(Batch_Size, device=device) - 0.5) * 6.0
|
||||
|
||||
a_t, a_c = gen_a_batch(theta_t), gen_a_batch(theta_c)
|
||||
|
||||
# 构建纯净的基准协方差矩阵
|
||||
R_interf = noise_power * torch.eye(N_ant, device=device).unsqueeze(0) + INR_linear * (a_c @ a_c.mH)
|
||||
w_opt = torch.linalg.pinv(R_interf) @ a_t
|
||||
w_opt = w_opt / torch.norm(w_opt, dim=1, keepdim=True)
|
||||
|
||||
F_RF = model(w_opt)
|
||||
F_BB = torch.linalg.pinv(F_RF.mH @ F_RF + 1e-3*torch.eye(N_RF, device=device).unsqueeze(0)) @ (F_RF.mH @ w_opt)
|
||||
w_ai = F_RF @ F_BB
|
||||
w_ai = w_ai / torch.norm(w_ai, dim=1, keepdim=True)
|
||||
|
||||
# 波束相似度
|
||||
sim = torch.abs(torch.sum(w_opt.conj() * w_ai, dim=(1,2)))**2
|
||||
loss_sim = torch.mean(1.0 - sim)
|
||||
# 杂波方向绝对抑制
|
||||
loss_null = torch.mean(torch.abs(torch.sum(a_c.conj() * w_ai, dim=(1,2)))**2)
|
||||
|
||||
# 【重拳出击】:给予杂波方向 150 倍的极限死亡惩罚!逼迫 AI 挖出深坑!
|
||||
loss = loss_sim + 150.0 * loss_null
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
loss_history.append(loss.item())
|
||||
|
||||
print(f"训练完成!耗时: {time.time() - start_time:.2f} 秒\n")
|
||||
|
||||
# ==========================================
|
||||
# 3. 静态成果检验:生成完美的波束方向图
|
||||
# ==========================================
|
||||
print(f"正在生成目标 {test_t_ang}°,强杂波 {test_c_ang}° 的终极抗干扰方向图...")
|
||||
theta_scan = torch.arange(-90, 90.1, 0.1, device=device)
|
||||
A_scan = gen_a_batch(theta_scan).squeeze(2).T
|
||||
|
||||
a_t_stat = gen_a_batch(torch.tensor([test_t_ang], device=device))
|
||||
a_c_stat = gen_a_batch(torch.tensor([test_c_ang], device=device))
|
||||
R_interf_stat = noise_power * torch.eye(N_ant, device=device).unsqueeze(0) + INR_linear * (a_c_stat @ a_c_stat.mH)
|
||||
|
||||
with torch.no_grad():
|
||||
model.eval()
|
||||
w_opt_stat = torch.linalg.pinv(R_interf_stat) @ a_t_stat
|
||||
w_opt_stat /= torch.norm(w_opt_stat, dim=1, keepdim=True)
|
||||
|
||||
F_RF_stat = model(w_opt_stat)
|
||||
F_BB_stat = torch.linalg.pinv(F_RF_stat.mH @ F_RF_stat + 1e-3*torch.eye(N_RF, device=device).unsqueeze(0)) @ (F_RF_stat.mH @ w_opt_stat)
|
||||
w_ai_stat = F_RF_stat @ F_BB_stat
|
||||
w_ai_stat /= torch.norm(w_ai_stat, dim=1, keepdim=True)
|
||||
|
||||
# 极简且精准的点乘测向
|
||||
gain_dig = torch.abs(torch.sum(w_opt_stat.conj() * A_scan.unsqueeze(0), dim=1))**2
|
||||
BP_dig = 10 * torch.log10(gain_dig).cpu().numpy().flatten()
|
||||
BP_dig -= np.max(BP_dig)
|
||||
|
||||
gain_ai = torch.abs(torch.sum(w_ai_stat.conj() * A_scan.unsqueeze(0), dim=1))**2
|
||||
BP_ai = 10 * torch.log10(gain_ai).cpu().numpy().flatten()
|
||||
BP_ai -= np.max(BP_ai)
|
||||
|
||||
# ==========================================
|
||||
# 4. 绘制终极展示双图
|
||||
# ==========================================
|
||||
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(14, 5))
|
||||
|
||||
ax1.plot(loss_history, color='#0072BD', linewidth=2)
|
||||
ax1.set_title('扇区聚焦网络 (TSC-DL) 极限优化损失曲线', fontweight='bold', fontsize=12)
|
||||
ax1.set_xlabel('迭代轮数 (Epochs)')
|
||||
ax1.set_ylabel('联合 Loss')
|
||||
ax1.grid(True, ls='--')
|
||||
|
||||
ax2.plot(theta_scan.cpu().numpy(), BP_dig, color='#000000', linestyle='-', linewidth=2, alpha=0.5, label='纯数字理想波束')
|
||||
ax2.plot(theta_scan.cpu().numpy(), BP_ai, color='#D95319', linestyle='--', linewidth=2.5, label='扇区聚焦 AI 混合波束')
|
||||
ax2.axvline(test_t_ang, color='g', linestyle='-.', label=f'机动目标 ({test_t_ang}°)')
|
||||
ax2.axvline(test_c_ang, color='r', linestyle='-.', label=f'强杂波 ({test_c_ang}°)')
|
||||
ax2.set_ylim([-60, 0])
|
||||
ax2.set_xlim([-60, 60])
|
||||
ax2.set_title(f'深度学习端到端波束响应 (目标 {test_t_ang}° 杂波 {test_c_ang}° )', fontweight='bold', fontsize=12)
|
||||
ax2.set_xlabel('角度 θ (°)')
|
||||
ax2.set_ylabel('归一化增益 (dB)')
|
||||
ax2.grid(True, ls='--')
|
||||
ax2.legend(loc='lower center', fontsize=10)
|
||||
|
||||
plt.tight_layout()
|
||||
plt.savefig('DL_Sector_Focus_Final.png', dpi=300)
|
||||
plt.show()
|
||||
Reference in New Issue
Block a user