{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "4aac8cae-7710-4f50-9ccf-06ee986885f1",
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import time\n",
    "import cv2\n",
    "import numpy as np\n",
    "import vision.utils.box_utils_numpy as box_utils\n",
    "import onnxruntime as ort\n",
    "\n",
    "# 定义预测函数，对模型输出的边界框和置信度进行后处理\n",
    "def predict(width, height, confidences, boxes, prob_threshold, iou_threshold=0.3, top_k=-1):\n",
    "    boxes = boxes[0]\n",
    "    confidences = confidences[0]\n",
    "    picked_box_probs = []\n",
    "    picked_labels = []\n",
    "    for class_index in range(1, confidences.shape[1]):\n",
    "        probs = confidences[:, class_index]\n",
    "        mask = probs > prob_threshold\n",
    "        probs = probs[mask]\n",
    "        if probs.shape[0] == 0:\n",
    "            continue\n",
    "        subset_boxes = boxes[mask, :]\n",
    "        box_probs = np.concatenate([subset_boxes, probs.reshape(-1, 1)], axis=1)\n",
    "        box_probs = box_utils.hard_nms(box_probs,\n",
    "                                       iou_threshold=iou_threshold,\n",
    "                                       top_k=top_k,\n",
    "                                       )\n",
    "        picked_box_probs.append(box_probs)\n",
    "        picked_labels.extend([class_index] * box_probs.shape[0])\n",
    "    if not picked_box_probs:\n",
    "        return np.array([]), np.array([]), np.array([])\n",
    "    picked_box_probs = np.concatenate(picked_box_probs)\n",
    "    picked_box_probs[:, 0] *= width\n",
    "    picked_box_probs[:, 1] *= height\n",
    "    picked_box_probs[:, 2] *= width\n",
    "    picked_box_probs[:, 3] *= height\n",
    "    return picked_box_probs[:, :4].astype(np.int32), np.array(picked_labels), picked_box_probs[:, 4]\n",
    "\n",
    "# 从标签文件中读取每一行，并去除行首尾的空白字符，得到类别名称列表 2分\n",
    "class_names = [_______________ for name in open('voc-model-labels.txt').readlines()]\n",
    "\n",
    "# 创建 ONNX Runtime 的推理会话，用于运行模型进行推理 2分\n",
    "ort_session = _______________('version-RFB-320.onnx')\n",
    "\n",
    "# 获取模型输入的名称 2分\n",
    "input_name = _______________()[0].name\n",
    "\n",
    "# 定义保存检测结果图像的目录路径\n",
    "result_path = \"./detect_imgs_results_onnx\"\n",
    "\n",
    "# 定义置信度阈值，用于筛选出置信度较高的检测结果\n",
    "threshold = 0.7\n",
    "# 定义存储待检测图像的目录路径\n",
    "path = \"imgs\"\n",
    "# 用于统计所有图像中检测到的目标框总数，初始化为 0\n",
    "sum = 0\n",
    "\n",
    "# 如果保存结果的目录不存在，则创建该目录 2分\n",
    "if not os.path.exists(result_path):\n",
    "    os._______________\n",
    "    \n",
    "# 获取指定目录下的所有文件和文件夹名称列表\n",
    "listdir = os.listdir(path)\n",
    "\n",
    "# 遍历目录下的每个文件\n",
    "for file_path in listdir:\n",
    "    # 拼接图像文件的完整路径\n",
    "    img_path = os.path.join(path, file_path)\n",
    "    # 使用 OpenCV 读取图像文件 2分\n",
    "    orig_image = _______________\n",
    "    # 将图像从 BGR 颜色空间转换为 RGB 颜色空间（许多模型要求输入为 RGB 格式）\n",
    "    image = cv2.cvtColor(orig_image, cv2.COLOR_BGR2RGB)\n",
    "    # 将图像调整为 320x240 的尺寸（符合模型输入的尺寸要求） 2分\n",
    "    image = _______________(_______________, (320, 240))\n",
    "    # 定义图像归一化的均值数组 2分\n",
    "    image_mean = _______________([127, 127, 127])\n",
    "    # 对图像进行归一化处理，减去均值并除以 128\n",
    "    image = (image - image_mean) / 128\n",
    "    # 将图像的维度从 (高度, 宽度, 通道数) 转换为 (通道数, 高度, 宽度)\n",
    "    image = np.transpose(image, [2, 0, 1])\n",
    "    # 在第一个维度上扩展一个维度，将图像变为 (1, 通道数, 高度, 宽度)，以符合模型输入的维度要求  1分\n",
    "    image = _______________(image, axis=0)\n",
    "    # 将图像数据类型转换为 float32 类型\n",
    "    image = image.astype(np.float32)\n",
    "    # 记录开始时间，用于计算模型推理的耗时\n",
    "    time_time = time.time()\n",
    "    # 使用 ONNX Runtime 运行模型，输入图像数据，得到模型输出的置信度和边界框  2分\n",
    "    confidences, boxes = _______________(None, {input_name: image})\n",
    "    # 计算并打印模型推理的耗时\n",
    "    print(\"cost time:{}\".format(time.time() - time_time))\n",
    "    # 调用 predict 函数对模型输出的边界框和置信度进行后处理，得到最终的边界框、类别标签和置信度\n",
    "    boxes, labels, probs = predict(orig_image.shape[1], orig_image.shape[0], confidences, boxes, threshold)\n",
    "    # 遍历每个检测到的目标框\n",
    "    for i in range(boxes.shape[0]):\n",
    "        # 获取当前目标框的坐标\n",
    "        box = boxes[i, :]\n",
    "        # 生成当前目标框的标签字符串，包含类别名称和置信度\n",
    "        label = f\"{class_names[labels[i]]}: {probs[i]:.2f}\"\n",
    "\n",
    "        # 在原始图像上绘制目标框，颜色为 (255, 255, 0)，线条粗细为 4\n",
    "        cv2.rectangle(orig_image, (box[0], box[1]), (box[2], box[3]), (255, 255, 0), 4)\n",
    "        # 将绘制了目标框的图像保存到结果目录中\n",
    "        cv2.imwrite(os.path.join(result_path, file_path), orig_image)\n",
    "    # 累加当前图像中检测到的目标框数量到总数中\n",
    "    sum += boxes.shape[0]\n",
    "# 打印所有图像中检测到的目标框总数\n",
    "print(\"sum:{}\".format(sum))"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3 (ipykernel)",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.11.7"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
