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)
Logo

腾讯云面向开发者汇聚海量精品云计算使用和开发经验,营造开放的云计算技术生态圈。

更多推荐