如何使用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目录下找到对应的配置文件,里面会标注权重的获取方式,也可以用官方脚本自动下载,或者手动下载后放到你方便调用的路径。
第四步:加载模型做推理(直接用预训练模型)
如果你只是想先用预训练模型跑个测试,写几行代码就能搞定:
- 先导入需要的库:
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
- 加载配置和预训练权重:
# 替换成你本地的配置文件路径 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()
- 定义推理函数:
@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
- 加载图片并可视化结果:
# 加载测试图片 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
相关产品推荐
相关产品推荐

