1.准备配置文件:首先,你需要下载YOLO项目中的基础配置文件,以确保追踪功能正常工作。请激活你的yolov8环境,并用以下命令初始化Ultralytics的配置文件:

# 这会在本地生成一个 'ultralytics' 文件夹


yolo cfg 

执行完毕后,你会得到一个 ultralytics/cfg/default.yaml 文件,这是之后脚本读取配置的基础。

2.下载追踪器配置:YOLO官方提供了一个高度适配的目标追踪器配置。下载后,将其放在与Python脚本相同的目录下。

https://github.com/ultralytics/ultralytics/blob/main/ultralytics/cfg/trackers/bytetrack.yaml

3.完整python代码:

from ultralytics import YOLO
import cv2

# ========== 配置参数 ==========
MODEL_PATH = "runs/detect/runs/fruit_final/102img_150e/weights/best.pt"
SOURCE_VIDEO = r"D:\A-果实\果实视频\果实视频1.MP4"
TARGET_VIDEO = "processed_fruits_counting_video.mp4"
TRACKER_CONFIG = "bytetrack.yaml"   # 若没有该文件可改为 None
CONF_THRESHOLD = 0.5
IOU_THRESHOLD = 0.5


def main():
    print("[INFO] 加载模型...")
    model = YOLO(MODEL_PATH)

    cap = cv2.VideoCapture(SOURCE_VIDEO)
    if not cap.isOpened():
        print(f"[ERROR] 无法打开视频文件: {SOURCE_VIDEO}")
        return

    fps = int(cap.get(cv2.CAP_PROP_FPS))
    width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
    height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
    out = cv2.VideoWriter(TARGET_VIDEO, cv2.VideoWriter_fourcc(*'mp4v'), fps, (width, height))

    fruit_ids = set()
    frame_count = 0

    print("[INFO] 开始处理视频,按 'q' 键停止...")
    while cap.isOpened():
        success, frame = cap.read()
        if not success:
            break
        frame_count += 1

        # 追踪
        if TRACKER_CONFIG:
            results = model.track(frame, persist=True, conf=CONF_THRESHOLD,
                                  iou=IOU_THRESHOLD, tracker=TRACKER_CONFIG)
        else:
            results = model.track(frame, persist=True, conf=CONF_THRESHOLD, iou=IOU_THRESHOLD)

        annotated_frame = results[0].plot()

        boxes = results[0].boxes
        if boxes is not None and boxes.id is not None:
            track_ids = boxes.id.int().cpu().tolist()
            fruit_ids.update(track_ids)

        total_count = len(fruit_ids)
        cv2.putText(annotated_frame, f"Total Fruits: {total_count}", (10, 30),
                    cv2.FONT_HERSHEY_SIMPLEX, 1, (0, 255, 0), 2)
        cv2.putText(annotated_frame, f"Frame: {frame_count}", (10, 70),
                    cv2.FONT_HERSHEY_SIMPLEX, 0.7, (0, 255, 0), 2)

        out.write(annotated_frame)

        # ========== 实时显示 ==========
        cv2.imshow("Fruit Tracking", annotated_frame)
        if cv2.waitKey(1) & 0xFF == ord('q'):
            break

        if frame_count % 100 == 0:
            print(f"已处理 {frame_count} 帧, 当前累计果子数: {total_count}")

    cap.release()
    out.release()
    cv2.destroyAllWindows()
    print(f"\n[完成] 视频处理完毕!")
    print(f"最终统计结果: 视频中共检测到 {len(fruit_ids)} 个独立果子。")
    print(f"带标注的视频已保存至: {TARGET_VIDEO}")


if __name__ == "__main__":
    main()

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐