从OpenCV到YOLOv5,手把手整合完整视觉项目完整代码
·
import cv2
import torch
import numpy as np
import time
# -------------------------- 请修改这3个参数(必改)--------------------------
model_path = "best.pt" # YOLOv5模型路径(与代码同一目录填"best.pt")
video_path = 0 # 0=默认摄像头,可替换为视频路径(如:"test.mp4")
save_video = True # 是否保存检测结果视频(True=保存,False=不保存)
# --------------------------------------------------------------------------
# -------------------------- 可选参数(按需修改)--------------------------
conf_threshold = 0.4 # 置信度阈值(3.31优化技巧)
track_distance = 50 # 追踪距离阈值(避免重复计数)
alarm_line = (200, 350, 1080, 350) # 越线报警线(x1,y1,x2,y2)
# --------------------------------------------------------------------------
# 1. 初始化所有模块(整合本周所有核心功能)
# 1.1 加载YOLOv5模型(优化后,兼容所有版本)
model = torch.hub.load('ultralytics/yolov5', 'custom', path=model_path, force_reload=True)
model.conf = conf_threshold # 应用优化技巧,过滤无效框
class_names = ['cat', 'dog'] # 类别的,与数据集一致
colors = [(255, 0, 0), (0, 0, 255)] # 颜色(猫=蓝,狗=红)
# 1.2 初始化视频读取与保存(4.2进阶功能)
cap = cv2.VideoCapture(video_path)
cap.set(cv2.CAP_PROP_FRAME_WIDTH, 1280)
cap.set(cv2.CAP_PROP_FRAME_HEIGHT, 720)
# 初始化视频写入器
fourcc = cv2.VideoWriter_fourcc(*'mp4v')
fps = cap.get(cv2.CAP_PROP_FPS)
frame_width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
frame_height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
out = None
if save_video:
out = cv2.VideoWriter('complete_project_result.mp4', fourcc, fps, (frame_width, frame_height))
# 1.3 初始化计数、追踪、报警相关变量(4.1+4.2功能整合)
total_cat = 0 # 总猫数
total_dog = 0 # 总狗数
current_cat = 0 # 当前帧猫数
current_dog = 0 # 当前帧狗数
tracks = {} # 目标追踪字典
track_id = 0 # 目标唯一ID
crossed_ids = set() # 已越线目标ID(避免重复报警)
alarm_flag = False # 报警标志
# 2. 项目核心循环(整合所有功能,一步到位)
while cap.isOpened():
ret, frame = cap.read()
if not ret:
break # 视频读取完毕,退出循环
# 重置当前帧计数
current_cat = 0
current_dog = 0
# 3. 模块1:OpenCV图像预处理(3.27轮廓检测+画面优化)
# 灰度处理+轮廓提取(可选,可按需开启)
gray = cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY)
blur = cv2.GaussianBlur(gray, (5, 5), 0) # 模糊去噪
ret, thresh = cv2.threshold(blur, 127, 255, cv2.THRESH_BINARY)
contours, _ = cv2.findContours(thresh, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
# 绘制轮廓(可选,可视化预处理效果)
cv2.drawContours(frame, contours, -1, (0, 255, 0), 1)
# 4. 模块2:YOLOv5目标检测(3.30+3.31优化功能)
results = model(frame)
detections = results.pandas().xyxy[0].values # 提取检测结果
# 5. 模块3:目标追踪与计数(4.1基础功能)
current_tracks = []
for det in detections:
xmin, ymin, xmax, ymax, conf, cls, name = det
center_x = int((xmin + xmax) / 2)
center_y = int((ymin + ymax) / 2)
cls = int(cls)
# 当前帧计数
if cls == 0:
current_cat += 1
else:
current_dog += 1
# 目标追踪逻辑(优化后,更稳定)
matched = False
for track_id_exist, (track_center_x, track_center_y, track_cls) in tracks.items():
distance = np.sqrt((center_x - track_center_x)**2 + (center_y - track_center_y)**2)
if distance < track_distance and track_cls == cls:
tracks[track_id_exist] = (center_x, center_y, cls)
current_tracks.append(track_id_exist)
matched = True
# 模块4:越线报警(4.2进阶功能)
line_y = alarm_line[1]
if track_center_y > line_y and track_id_exist not in crossed_ids:
crossed_ids.add(track_id_exist)
alarm_flag = True
# 报警可视化+控制台提示
cv2.putText(frame, f"ALARM: {name} crossed line!", (500, 40),
cv2.FONT_HERSHEY_SIMPLEX, 1, (0, 0, 255), 3)
print(f"【报警】{name} 越过报警线,时间:{time.strftime('%H:%M:%S')}")
break
# 新目标计数
if not matched:
tracks[track_id] = (center_x, center_y, cls)
current_tracks.append(track_id)
if cls == 0:
total_cat += 1
else:
total_dog += 1
track_id += 1
# 删除消失的目标轨迹
tracks = {k: v for k, v in tracks.items() if k in current_tracks}
# 6. 模块5:可视化展示(整合所有可视化元素)
# 绘制报警线
cv2.line(frame, (alarm_line[0], alarm_line[1]), (alarm_line[2], alarm_line[3]), (0, 255, 0), 2)
cv2.putText(frame, "Alarm Line", (alarm_line[0], alarm_line[1]-10),
cv2.FONT_HERSHEY_SIMPLEX, 0.5, (0, 255, 0), 2)
# 绘制检测框、类别、置信度
for det in detections:
xmin, ymin, xmax, ymax, conf, cls, name = det
xmin, ymin, xmax, ymax = int(xmin), int(ymin), int(xmax), int(ymax)
cls = int(cls)
cv2.rectangle(frame, (xmin, ymin), (xmax, ymax), colors[cls], 2)
label = f"{name} {conf:.2f}"
cv2.putText(frame, label, (xmin, ymin-10), cv2.FONT_HERSHEY_SIMPLEX, 0.5, colors[cls], 2)
# 绘制目标轨迹
for track_id_exist, (center_x, center_y, cls) in tracks.items():
cv2.circle(frame, (center_x, center_y), 5, colors[cls], -1)
if track_id_exist > 0:
prev_center = tracks.get(track_id_exist - 1, None)
if prev_center and prev_center[2] == cls:
cv2.line(frame, (prev_center[0], prev_center[1]), (center_x, center_y), colors[cls], 2)
# 绘制计数信息(总计数+当前帧计数)
cv2.putText(frame, f"Total Cat: {total_cat}", (20, 40), cv2.FONT_HERSHEY_SIMPLEX, 1, (255, 0, 0), 2)
cv2.putText(frame, f"Total Dog: {total_dog}", (20, 80), cv2.FONT_HERSHEY_SIMPLEX, 1, (0, 0, 255), 2)
cv2.putText(frame, f"Current Cat: {current_cat}", (1000, 40), cv2.FONT_HERSHEY_SIMPLEX, 1, (255, 0, 0), 2)
cv2.putText(frame, f"Current Dog: {current_dog}", (1000, 80), cv2.FONT_HERSHEY_SIMPLEX, 1, (0, 0, 255), 2)
# 绘制报警提示(持续3秒)
if alarm_flag:
cv2.putText(frame, "ALARM! Target Crossed Line", (400, 700),
cv2.FONT_HERSHEY_SIMPLEX, 1.5, (0, 0, 255), 3)
alarm_flag = False
# 保存检测结果视频(模块4功能)
if save_video and out is not None:
out.write(frame)
# 显示完整项目画面
cv2.imshow("Complete Vision Project (OpenCV + YOLOv5)", frame)
# 按「q」键退出
if cv2.waitKey(1) & 0xFF == ord('q'):
break
# 释放所有资源(避免内存泄漏)
cap.release()
if save_video and out is not None:
out.release()
cv2.destroyAllWindows()
# 打印最终项目统计结果(便于汇报、整理数据)
print("="*60)
print("完整视觉项目 - 最终统计结果")
print("="*60)
print(f"总检测到猫:{total_cat} 只")
print(f"总检测到狗:{total_dog} 只")
print(f"越线目标数量:{len(crossed_ids)} 个")
print(f"检测结果视频:{'已保存(complete_project_result.mp4)' if save_video else '未保存'}")
print(f"项目运行状态:成功完成")
print("="*60)
更多推荐
所有评论(0)