Skip to content
Closed
Changes from all commits
Commits
Show all changes
44 commits
Select commit Hold shift + click to select a range
a879469
完成选题毕进行首次提交
MKHQ07 Sep 22, 2025
5c3e0bf
完成选题毕进行首次提交
MKHQ07 Sep 22, 2025
c8cefda
Merge remote-tracking branch 'origin/main'
MKHQ07 Sep 22, 2025
290feda
Merge branch 'OpenHUTB:main' into main
MKHQ07 Sep 23, 2025
d478785
Merge branch 'OpenHUTB:main' into main
MKHQ07 Sep 29, 2025
8c2818c
Merge branch 'OpenHUTB:main' into main
MKHQ07 Oct 13, 2025
2d1a00b
Merge branch 'OpenHUTB:main' into main
MKHQ07 Oct 20, 2025
4ecca9a
基于深度学习的无人机控制与可视化系统
MKHQ07 Oct 20, 2025
c8e7221
基于深度学习的无人机控制与可视化系统
MKHQ07 Oct 20, 2025
3e68432
Merge branch 'OpenHUTB:main' into main
MKHQ07 Oct 27, 2025
3ad4ea0
基于深度学习的无人机控制与可视化系统
MKHQ07 Oct 27, 2025
fbc0120
Merge branch 'OpenHUTB:main' into main
MKHQ07 Oct 27, 2025
00e8ed2
基于深度学习的无人机控制与可视化系统
MKHQ07 Oct 27, 2025
1074cab
基于深度学习的无人机控制与可视化系统
MKHQ07 Oct 27, 2025
c1f19fd
目标跟踪
MKHQ07 Oct 27, 2025
4ba6c94
Merge remote-tracking branch 'origin/main'
MKHQ07 Oct 27, 2025
68d026c
Merge branch 'OpenHUTB:main' into main
MKHQ07 Oct 29, 2025
bb5390f
Merge branch 'OpenHUTB:main' into main
MKHQ07 Nov 3, 2025
2452e35
高度保持
MKHQ07 Nov 3, 2025
9c662ab
Merge remote-tracking branch 'origin/main'
MKHQ07 Nov 3, 2025
ee2e73a
Merge branch 'OpenHUTB:main' into main
MKHQ07 Nov 3, 2025
7983cb4
速度保持
MKHQ07 Nov 3, 2025
ed33ebd
速度保持
MKHQ07 Nov 3, 2025
a802f3e
速度保持
MKHQ07 Nov 3, 2025
8785c8f
速度保持
MKHQ07 Nov 3, 2025
d251672
速度保持
MKHQ07 Nov 3, 2025
7010c26
速度保持
MKHQ07 Nov 3, 2025
7387c85
速度保持.
MKHQ07 Nov 3, 2025
8a8c731
Merge remote-tracking branch 'origin/main'
MKHQ07 Nov 3, 2025
7bd1833
Merge branch 'OpenHUTB:main' into main
MKHQ07 Nov 3, 2025
f3e1d00
Merge branch 'OpenHUTB:main' into main
MKHQ07 Nov 17, 2025
788c942
平衡保持
MKHQ07 Nov 17, 2025
073a498
平衡保持
MKHQ07 Nov 17, 2025
4241a3a
平衡保持
MKHQ07 Dec 1, 2025
47234f8
图像分类
MKHQ07 Dec 15, 2025
482a44f
图像分类
MKHQ07 Dec 15, 2025
e485b57
图像分类
MKHQ07 Dec 15, 2025
ee252a4
图像分类
MKHQ07 Dec 15, 2025
d52716a
Merge branch 'OpenHUTB:main' into main
MKHQ07 Dec 15, 2025
ea0312f
图像分类
MKHQ07 Dec 15, 2025
d50d84f
图像分类
MKHQ07 Dec 15, 2025
8ab9541
Merge branch 'OpenHUTB:main' into main
MKHQ07 Dec 16, 2025
db6f483
图像分类
MKHQ07 Dec 16, 2025
3fec6f6
Merge remote-tracking branch 'origin/main'
MKHQ07 Dec 16, 2025
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
273 changes: 163 additions & 110 deletions src/driverless_car/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,40 +5,38 @@
from torchvision import datasets, transforms
import matplotlib.pyplot as plt
import numpy as np
import time


# 修复1:解决中文字体问题(改用英文显示,避免字体依赖)
# 若需中文,可注释以下行并使用系统自带中文字体,如:
# plt.rcParams["font.family"] = ["Microsoft YaHei", "SimSun", "Arial"]
# plt.rcParams["axes.unicode_minus"] = False
# --------------------------
# 直接使用英文显示类别,避免中文字体问题
classes = ('Airplane', 'Car', 'Bird', 'Cat', 'Deer', 'Dog', 'Frog', 'Horse', 'Ship', 'Truck')
import os
from PIL import Image
import random

# --------------------------
# 修复2:设置Matplotlib后端为TkAgg,解决tostring_rgb报错
# 1. 基础配置(解决可视化和字体问题)
# --------------------------
import matplotlib
matplotlib.use('TkAgg') # 更换后端,兼容PyCharm的可视化

# 类别名称(对应CIFAR-10的10类)
classes = ('Airplane', 'Car', 'Bird', 'Cat', 'Deer', 'Dog', 'Frog', 'Horse', 'Ship', 'Truck')

# --------------------------
# 1. 数据预处理与加载(模拟无人机采集的图像数据
# 2. 数据预处理(简化,减少计算量
# --------------------------
transform = transforms.Compose([
# 简化变换:移除数据增强(加快训练,牺牲一点泛化能力)
basic_transform = transforms.Compose([
transforms.Resize((32, 32)),
transforms.ToTensor(),
transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
])

train_dataset = datasets.CIFAR10(root='./data', train=True, download=True, transform=transform)
test_dataset = datasets.CIFAR10(root='./data', train=False, download=True, transform=transform)

train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=64, shuffle=False)
# 加载CIFAR-10数据集(使用简化变换)
train_dataset = datasets.CIFAR10(root='./data', train=True, download=True, transform=basic_transform)
test_dataset = datasets.CIFAR10(root='./data', train=False, download=True, transform=basic_transform)
# 增大批次大小,加快训练(根据内存调整,默认128)
train_loader = DataLoader(train_dataset, batch_size=128, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=128, shuffle=False)

# --------------------------
# 2. 搭建轻量化CNN模型(无人机端深度学习模型
# 3. 搭建轻量化CNN模型(原模型,参数少,速度快
# --------------------------
class DroneCNN(nn.Module):
def __init__(self):
Expand All @@ -50,7 +48,7 @@ def __init__(self):
self.fc1 = nn.Linear(128 * 4 * 4, 512)
self.fc2 = nn.Linear(512, 10)
self.relu = nn.ReLU()
self.dropout = nn.Dropout(0.5)
self.dropout = nn.Dropout(0.3) # 降低dropout比例,加快计算

def forward(self, x):
x = self.pool(self.relu(self.conv1(x)))
Expand All @@ -62,148 +60,203 @@ def forward(self, x):
x = self.fc2(x)
return x

# 初始化模型、损失函数、优化器
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = DroneCNN().to(device)
criterion = nn.CrossEntropyLoss()
# 使用SGD优化器(比Adam稍快,或保留Adam,影响不大)
optimizer = optim.Adam(model.parameters(), lr=0.001)

# --------------------------
# 3. 可视化函数(修复非阻塞显示问题
# 4. 训练函数(大幅优化速度
# --------------------------
def show_dataset_samples():
"""显示数据集的样本图像(模拟无人机采集的图像)"""
data_iter = iter(train_loader)
images, labels = next(data_iter)
images = images / 2 + 0.5 # 反归一化

plt.figure(figsize=(10, 6))
for i in range(12):
plt.subplot(3, 4, i+1)
plt.imshow(np.transpose(images[i].numpy(), (1, 2, 0)))
plt.title(classes[labels[i]])
plt.axis('off')
plt.suptitle('Drone Collected Image Samples (CIFAR-10 Simulation)', fontsize=14)
plt.tight_layout()
plt.show(block=False) # 保留非阻塞
plt.pause(0.1) # 修复:添加pause,解决后端渲染问题
# 移除time.sleep,改用plt.pause更稳定
# 全局开启交互模式,用于实时绘图
plt.ion()
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 4)) # 训练曲线窗口

def plot_training_curve(train_losses, train_accs, test_accs):
"""绘制训练损失和准确率曲线"""
plt.figure(figsize=(12, 4))
"""更新训练曲线窗口,减少绘图开销"""
ax1.clear()
ax2.clear()
# 损失曲线
plt.subplot(1, 2, 1)
plt.plot(train_losses, label='Training Loss')
plt.xlabel('Iteration (Batch)')
plt.ylabel('Loss')
plt.title('Training Loss Change')
plt.legend()
plt.grid(True)
ax1.plot(train_losses, label='Training Loss', color='blue')
ax1.set_xlabel('Iteration (Batch)')
ax1.set_ylabel('Loss')
ax1.set_title('Training Loss Change')
ax1.legend()
ax1.grid(True)
# 准确率曲线
plt.subplot(1, 2, 2)
plt.plot(train_accs, label='Training Accuracy')
plt.plot(test_accs, label='Test Accuracy')
plt.xlabel('Iteration (Batch)')
plt.ylabel('Accuracy (%)')
plt.title('Training/Test Accuracy Change')
plt.legend()
plt.grid(True)
plt.tight_layout()
plt.show(block=False)
plt.pause(0.1) # 修复:添加pause
ax2.plot(train_accs, label='Training Accuracy', color='green')
ax2.plot(test_accs, label='Test Accuracy', color='red')
ax2.set_xlabel('Iteration (Batch)')
ax2.set_ylabel('Accuracy (%)')
ax2.set_title('Training/Test Accuracy Change')
ax2.legend()
ax2.grid(True)
plt.draw()
plt.pause(0.01) # 减少暂停时间,加快绘图

def show_predictions():
"""显示模型的预测结果"""
def calculate_test_acc_fast(test_loader, sample_batches=5):
"""快速计算测试集准确率:只抽取少量批次,不遍历整个测试集"""
test_correct = 0
test_total = 0
model.eval()
data_iter = iter(test_loader)
images, labels = next(data_iter)
images = images.to(device)
labels = labels.to(device)

outputs = model(images)
_, predicted = torch.max(outputs, 1)
images = images / 2 + 0.5 # 反归一化

plt.figure(figsize=(12, 8))
for i in range(16):
plt.subplot(4, 4, i+1)
plt.imshow(np.transpose(images[i].cpu().numpy(), (1, 2, 0)))
true_label = classes[labels[i]]
pred_label = classes[predicted[i]]
color = 'green' if true_label == pred_label else 'red'
plt.title(f'True: {true_label}\nPred: {pred_label}', color=color)
plt.axis('off')
plt.suptitle('Drone Image Classification Model Predictions', fontsize=14)
plt.tight_layout()
plt.show(block=False)
plt.pause(0.1) # 修复:添加pause
with torch.no_grad():
# 只取前sample_batches个批次,大幅减少耗时
for i, (test_inputs, test_labels) in enumerate(test_loader):
if i >= sample_batches:
break
test_inputs, test_labels = test_inputs.to(device), test_labels.to(device)
test_outputs = model(test_inputs)
_, test_predicted = torch.max(test_outputs.data, 1)
test_total += test_labels.size(0)
test_correct += (test_predicted == test_labels).sum().item()
model.train()
if test_total == 0:
return 0.0
return 100 * test_correct / test_total

# --------------------------
# 4. 训练模型并实时可视化
# --------------------------
def train_model(epochs=2):
def train_model(epochs=1): # 减少训练轮数,默认1轮
train_losses = []
train_accs = []
test_accs = []
model.train()

# 显示数据集样本
show_dataset_samples()

for epoch in range(epochs):
running_loss = 0.0
correct = 0
total = 0
# 每200个批次更新一次曲线(原100,减少更新频率)
update_interval = 200
for i, (inputs, labels) in enumerate(train_loader):
inputs, labels = inputs.to(device), labels.to(device)

# 前向传播+反向传播+优化
optimizer.zero_grad()
outputs = model(inputs)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()

# 统计指标
running_loss += loss.item()
_, predicted = torch.max(outputs.data, 1)
total += labels.size(0)
correct += (predicted == labels).sum().item()

# 每100个批次记录并可视化
if i % 100 == 99:
train_loss = running_loss / 100
# 降低更新频率,减少计算和绘图开销
if i % update_interval == update_interval - 1:
train_loss = running_loss / update_interval
train_acc = 100 * correct / total
train_losses.append(train_loss)
train_accs.append(train_acc)

# 计算测试集准确率
test_correct = 0
test_total = 0
with torch.no_grad():
for test_inputs, test_labels in test_loader:
test_inputs, test_labels = test_inputs.to(device), test_labels.to(device)
test_outputs = model(test_inputs)
_, test_predicted = torch.max(test_outputs.data, 1)
test_total += test_labels.size(0)
test_correct += (test_predicted == test_labels).sum().item()
test_acc = 100 * test_correct / test_total
# 快速计算测试集准确率(只取5个批次)
test_acc = calculate_test_acc_fast(test_loader, sample_batches=5)
test_accs.append(test_acc)

print(f'Epoch {epoch+1}, Batch {i+1} | Loss: {train_loss:.3f} | Train Acc: {train_acc:.2f}% | Test Acc: {test_acc:.2f}%')
running_loss = 0.0
correct = 0
total = 0

# 实时绘制曲线
plot_training_curve(train_losses, train_accs, test_accs)

# 显示预测结果
show_predictions()
# 修复:最后添加plt.show(block=True),防止窗口闪退
# 保存模型
torch.save(model.state_dict(), 'drone_model.pth')
print('模型已保存为drone_model.pth')
model.eval()
plt.ioff() # 关闭交互模式
plt.show(block=False)

# --------------------------
# 5. 模拟无人机实时图像输入(优化推理速度)
# --------------------------
def load_drone_images(folder_path):
"""读取本地文件夹中的图像,模拟无人机采集的图像流"""
img_extensions = ['.jpg', '.jpeg', '.png', '.bmp', '.gif']
img_paths = []
for file in os.listdir(folder_path):
if os.path.splitext(file)[1].lower() in img_extensions:
img_paths.append(os.path.join(folder_path, file))
if not img_paths:
raise ValueError(f'文件夹{folder_path}中未找到任何图像文件!')
return img_paths

def preprocess_image(img_path):
"""预处理单张图像,简化操作"""
# 读取图像(RGB模式)
img = Image.open(img_path).convert('RGB')
original_img = img.copy()
# 预处理
img = basic_transform(img)
# 添加batch维度
img = torch.unsqueeze(img, 0)
return original_img, img.to(device)

def drone_real_time_inference(folder_path, delay=0.1): # 降低延迟,默认0.1秒
"""模拟无人机实时图像输入,优化推理速度"""
print(f'\n开始模拟无人机实时图像流(读取文件夹:{folder_path}),每{delay}秒处理一张图像...')
img_paths = load_drone_images(folder_path)
# 创建单个显示窗口,减少窗口创建开销
fig, ax = plt.subplots(figsize=(6, 4)) # 缩小窗口,加快绘图
plt.ion()

for img_path in img_paths:
try:
# 预处理图像
original_img, input_tensor = preprocess_image(img_path)

# 模型预测(简化,减少计算)
with torch.no_grad():
outputs = model(input_tensor)
probabilities = torch.softmax(outputs, dim=1)
pred_idx = torch.argmax(probabilities, dim=1).item()
pred_class = classes[pred_idx]
pred_conf = probabilities[0][pred_idx].item() * 100

# 可视化优化:只更新图像和文本,不重建窗口
ax.clear()
ax.imshow(original_img)
ax.axis('off')
# 简化文本显示,减少渲染开销
text = f'{pred_class} ({pred_conf:.1f}%)'
ax.text(5, 5, text, fontsize=10, color='red',
bbox=dict(facecolor='white', alpha=0.7, edgecolor='none'))
ax.set_title('Drone Real-Time View', fontsize=12)
plt.draw()
plt.pause(delay)

# 简化控制台输出
print(f'图像:{os.path.basename(img_path)} → {pred_class} ({pred_conf:.1f}%)')

except Exception as e:
print(f'处理图像{img_path}时出错:{e}')
continue

# 结束后保持窗口
plt.ioff()
ax.text(0.5, 0.5, 'Done!', fontsize=14, ha='center', va='center',
transform=ax.transAxes, bbox=dict(facecolor='red', alpha=0.8))
plt.draw()
plt.show(block=True)
print('Training Finished!')

# --------------------------
# 运行主程序
# 主程序运行
# --------------------------
if __name__ == '__main__':
train_model(epochs=2)
# 第一步:训练模型(优化后速度大幅提升)
train_model(epochs=1) # 可改为2轮,仍比原来快很多

# 第二步:加载模型(可选)
# model.load_state_dict(torch.load('drone_model.pth', map_location=device))
# model.eval()
# print('模型已加载')

# 第三步:模拟无人机实时图像输入
drone_image_folder = r"C:\Users\hyq52\Desktop\P1\potoh"

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

需要改成相对路径,否则运行其他机器运行不了

if not os.path.exists(drone_image_folder):
os.makedirs(drone_image_folder)
print(f'已创建文件夹:{drone_image_folder},请放入测试图片后重新运行!')
else:
drone_real_time_inference(drone_image_folder, delay=0.1) # 低延迟,快速播放
Loading