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

拥有2000+自定义农场结构训练图,如何用TensorFlow自定义模型?

嗨,我之前刚好做过针对自定义数据集的TensorFlow目标检测模型搭建,结合你这2000多张农场结构图像,给你梳理一套非常落地的具体步骤,比官方文档的宏观说明实用多了:

第一步:数据集预处理(核心基础,必须先搞定)
  • 标注你的农场结构目标:因为是自定义类别,得给每张图里的大棚、仓库、围栏这些农场结构做边界框标注。推荐用LabelImg工具,操作简单,导出Pascal VOC格式的XML文件即可。2000张图手动标注完全可行,记得按7:2:1的比例划分训练集、验证集、测试集。
  • 转成TensorFlow支持的TFRecord格式:TensorFlow Object Detection API首选TFRecord格式,官方在models/research/object_detection/dataset_tools路径下提供了create_pascal_tf_record.py脚本。你需要修改脚本里的标签映射(把你的农场结构类别对应成数字ID),然后运行脚本生成train.record和val.record。
  • 创建标签映射文件:写一个label_map.pbtxt文件,格式示例如下,把你所有的农场结构类别都列进去,ID要和TFRecord里的对应一致:
    item {
      id: 1
      name: 'greenhouse'
    }
    item {
      id: 2
      name: 'warehouse'
    }
    item {
      id: 3
      name: 'fence'
    }
    
第二步:选择并配置预训练基础模型(不用从零造轮子)
  • 挑选适配的预训练模型:TensorFlow Model Zoo里有大量预训练模型,比如SSD MobileNet(速度快,适合中小目标)、Faster R-CNN(精度高,适合复杂目标)、EfficientDet(平衡速度与精度)。下载对应模型的预训练权重(.tar.gz文件),解压后得到ckpt格式的权重文件。
  • 修改模型配置文件:在models/research/object_detection/configs里找到对应模型的.config文件(比如ssd_mobilenet_v2_fpnlite_320x320_coco17_tpu-8.config),重点修改以下内容:
    • 将num_classes改成你的农场结构类别总数;
    • 调整train_input_reader里的input_path指向你的train.record,label_map_path指向你的label_map.pbtxt;
    • 修改eval_input_reader里的input_path指向你的val.record,label_map_path同样指向label_map.pbtxt;
    • 设置fine_tune_checkpoint为你下载的预训练权重路径(比如ssd_mobilenet_v2_fpnlite_320x320_coco17_tpu-8/checkpoint/ckpt-0);
    • 根据你的GPU显存调整batch_size(比如16或32,显存不足就改小);
    • 设定训练步数,比如num_steps=20000,num_eval_steps=1000。
第三步:搭建训练环境并启动训练
  • 确认环境依赖:确保安装了TensorFlow 2.x版本,以及protobuf、pillow、lxml、matplotlib等依赖,用pip安装:
    pip install tensorflow protobuf pillow lxml matplotlib
    
    然后将models的核心路径加入Python环境变量:
    export PYTHONPATH=$PYTHONPATH:/path/to/models/research:/path/to/models/research/slim
    
    最后编译protobuf文件:
    cd /path/to/models/research
    protoc object_detection/protos/*.proto --python_out=.
    
  • 启动训练:用TF2版本的训练脚本model_main_tf2.py,命令示例如下:
    python /path/to/models/research/object_detection/model_main_tf2.py \
      --model_dir=/path/to/your/training_logs \
      --pipeline_config_path=/path/to/your/modified_config.config \
      --num_train_steps=20000 \
      --num_eval_steps=1000 \
      --alsologtostderr
    
    训练过程中可以用TensorBoard监控损失曲线:
    tensorboard --logdir=/path/to/your/training_logs
    
第四步:导出并测试推理模型
  • 导出可推理模型:训练完成后,用export_tflite_graph_tf2.py导出saved_model格式的推理模型:
    python /path/to/models/research/object_detection/export_tflite_graph_tf2.py \
      --pipeline_config_path=/path/to/your/modified_config.config \
      --trained_checkpoint_dir=/path/to/your/training_logs \
      --output_directory=/path/to/your/exported_model
    
    如果需要移动端部署,还可以进一步转成TFLite格式。
  • 测试模型效果:写个简单的Python脚本加载模型,测试农场结构检测:
    import tensorflow as tf
    from object_detection.utils import label_map_util
    from object_detection.utils import visualization_utils as viz_utils
    import cv2
    
    # 加载训练好的模型
    detect_fn = tf.saved_model.load('/path/to/your/exported_model/saved_model')
    # 加载标签映射
    category_index = label_map_util.create_category_index_from_labelmap('/path/to/your/label_map.pbtxt', use_display_name=True)
    
    # 读取测试图像
    image = cv2.imread('/path/to/your/test_farm_image.jpg')
    input_tensor = tf.convert_to_tensor(image)
    input_tensor = input_tensor[tf.newaxis, ...]
    
    # 执行推理
    detections = detect_fn(input_tensor)
    
    # 可视化检测结果
    viz_utils.visualize_boxes_and_labels_on_image_array(
        image,
        detections['detection_boxes'][0].numpy(),
        detections['detection_classes'][0].numpy().astype(int),
        detections['detection_scores'][0].numpy(),
        category_index,
        use_normalized_coordinates=True,
        max_boxes_to_draw=200,
        min_score_thresh=0.5,  # 过滤低置信度结果,可调整
        agnostic_mode=False)
    
    # 显示结果
    cv2.imshow('Farm Structure Detection', image)
    cv2.waitKey(0)
    
第五步:模型调优(可选但能大幅提升效果)
  • 精度不足时:尝试增大输入图像尺寸(比如把config里的image_resizer改成640x640)、换更大的预训练模型(比如EfficientDet D4)、增加训练步数、添加数据增强(在config的data_augmentation_options里加随机翻转、裁剪、亮度调整等)。
  • 速度不够时:选用更轻量的模型(比如SSD MobileNet V2)、减小输入尺寸、转TFLite时启用量化压缩。

内容的提问来源于stack exchange,提问作者Ajinkya

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.19 08:39:52