You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何使用TensorFlow Model Garden中的官方模型(以YOLOv7的TensorFlow版本为例)

如何使用TensorFlow Model Garden中的官方模型(以YOLOv7的TensorFlow版本为例)

嘿,我完全懂你的困惑——TensorFlow Model Garden确实放了不少官方模型,但上手的步骤有时候写得不够直白,尤其是像YOLOv7这类热门模型,刚接触的时候容易摸不着头脑。刚好我之前折腾过TF版本的YOLOv7,给你一步步捋清楚怎么用:

第一步:先把Model Garden的代码拿到本地

TensorFlow的很多模型(包括YOLOv7)依赖Model Garden里的自定义层、工具脚本,所以首先得把仓库代码弄下来。你可以用git克隆,或者直接下载压缩包解压,命令是:

git clone https://github.com/tensorflow/models.git

第二步:安装必要的依赖

进入models/research目录,先装基础依赖,再安装Object Detection API的包:

cd models/research
pip install -r requirements.txt
pip install .

这里要注意,如果遇到protobuf版本冲突的问题,试试装protobuf==3.20.0这个版本,很多时候能解决报错。

第三步:获取YOLOv7的预训练权重

Model Garden里的YOLOv7提供了基于COCO数据集训练好的预训练权重,你可以在object_detection/configs/yolov7目录下找到对应的配置文件,里面会标注权重的获取方式,也可以用官方脚本自动下载,或者手动下载后放到你方便调用的路径。

第四步:加载模型做推理(直接用预训练模型)

如果你只是想先用预训练模型跑个测试,写几行代码就能搞定:

  1. 先导入需要的库:
import tensorflow as tf
import cv2
import numpy as np
from object_detection.builders import model_builder
from object_detection.utils import config_util, visualization_utils as viz_utils
  1. 加载配置和预训练权重:
# 替换成你本地的配置文件路径
config_path = "models/research/object_detection/configs/yolov7/yolov7_coco_config.yaml"
# 替换成你的预训练权重路径
checkpoint_path = "path/to/your/yolov7_checkpoint"

# 加载模型配置
configs = config_util.get_configs_from_pipeline_file(config_path)
model_config = configs["model"]
# 构建检测模型(设置is_training=False表示用于推理)
detection_model = model_builder.build(model_config=model_config, is_training=False)
# 恢复预训练权重
ckpt = tf.compat.v2.train.Checkpoint(model=detection_model)
ckpt.restore(checkpoint_path).expect_partial()
  1. 定义推理函数:
@tf.function
def detect_fn(image):
    image, shapes = detection_model.preprocess(image)
    prediction_dict = detection_model.predict(image, shapes)
    detections = detection_model.postprocess(prediction_dict, shapes)
    return detections
  1. 加载图片并可视化结果:
# 加载测试图片
image_path = "test_image.jpg"
image = cv2.imread(image_path)
image_rgb = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
input_tensor = tf.convert_to_tensor(np.expand_dims(image_rgb, 0), dtype=tf.float32)

# 执行推理
detections = detect_fn(input_tensor)

# 处理推理结果
num_detections = int(detections.pop("num_detections"))
detections = {key: value[0, :num_detections].numpy() for key, value in detections.items()}
detections["num_detections"] = num_detections
detections["detection_classes"] = detections["detection_classes"].astype(np.int64)

# 加载COCO类别标签
label_map_path = "models/research/object_detection/data/mscoco_label_map.pbtxt"
category_index = viz_utils.create_category_index_from_labelmap(label_map_path, use_display_name=True)

# 在图片上画框和标签
viz_utils.visualize_boxes_and_labels_on_image_array(
    image_rgb,
    detections["detection_boxes"],
    detections["detection_classes"],
    detections["detection_scores"],
    category_index,
    use_normalized_coordinates=True,
    max_boxes_to_draw=200,
    min_score_thresh=0.3,
    agnostic_mode=False
)

# 显示结果
cv2.imshow("YOLOv7 Detection Result", cv2.cvtColor(image_rgb, cv2.COLOR_RGB2BGR))
cv2.waitKey(0)
cv2.destroyAllWindows()

第五步:用自己的数据集微调YOLOv7

如果要训练自己的数据集,得先把数据转换成TensorFlow Object Detection API支持的TFRecord格式,然后修改YOLOv7的配置文件:

  • 把配置里的train_input_reader和eval_input_reader路径改成你的TFRecord文件路径
  • 调整num_classes为你的数据集类别数
  • 根据需要修改训练步数、学习率等参数

然后用官方的训练脚本启动训练:

python models/research/object_detection/model_main_tf2.py \
  --model_dir=./yolov7_training_logs \
  --pipeline_config_path=models/research/object_detection/configs/yolov7/yolov7_coco_config.yaml

一些要注意的小坑

  • 确保你的TensorFlow版本和Model Garden要求匹配,建议用TF 2.8及以上的稳定版
  • 如果遇到导入报错,检查一下models/research目录有没有加到Python路径里,或者重新安装Object Detection API包
  • 预训练权重一定要和配置文件对应,比如用COCO配置就别拿其他数据集的权重

备注:内容来源于stack exchange,提问作者Only-A-User

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.22 12:40:28