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

如何基于TensorFlow Model Zoo预训练模型完成自定义数据集的目标检测迁移学习(非Keras实现,以Centernet_1024x1024为例)

Hey there! I’ve been working with TensorFlow’s Model Zoo for custom object detection without relying on Keras, so I can walk you through the exact process step by step. Here’s how to pull it off with a model like Centernet_1024x1024:

1. Prepare Your Custom Dataset

First, you’ll need to convert your dataset to TFRecord format—this is the standard input format for TensorFlow Object Detection API models (no Keras involved here).

  • Use the official conversion scripts from the Object Detection API repo:
    • For Pascal VOC-style datasets: create_pascal_tf_record.py
    • For COCO-style datasets: create_coco_tf_record.py
  • If your dataset uses a custom format, modify one of these scripts to parse your annotations and images into TFRecords.
  • Create a label_map.pbtxt file that maps class IDs to names, like this:
    item {
      id: 1
      name: 'your_class_1'
    }
    item {
      id: 2
      name: 'your_class_2'
    }
    
2. Download the Pre-trained Centernet Model

Grab the pre-trained checkpoint from TensorFlow Model Zoo:

  • Download the Centernet_1024x1024 package (look for the .tar.gz file containing the checkpoint)
  • Extract it to a directory (e.g., ./pretrained_centernet/)—you’ll need the ckpt-xxxx files inside for fine-tuning.
3. Configure the Pipeline Config File

Copy the default config file for your Centernet model (e.g., centernet_hg104_1024x1024_coco17_tpu-32.config from the Model Zoo or Object Detection API repo) and modify these critical sections:

  • num_classes: Set this to the number of classes in your custom dataset.
  • train_input_reader:
    • Update input_path to point to your training TFRecord file.
    • Update label_map_path to your label_map.pbtxt file.
  • eval_input_reader: Do the same as above for your validation TFRecord.
  • train_config:
    • Set fine_tune_checkpoint to the path of your pre-trained checkpoint (e.g., ./pretrained_centernet/ckpt-32768—note you don’t include the file extension).
    • Set fine_tune_checkpoint_type: "detection" (since we’re fine-tuning the full detection pipeline).
    • Adjust batch_size, learning_rate, and num_steps based on your hardware and dataset size.
4. Run Training with the Object Detection API

Use the official model_main_tf2.py script (this is the non-Keras, native TensorFlow training entry point). Run this command in your terminal:

python model_main_tf2.py \
  --model_dir=./training_output/ \
  --pipeline_config_path=./centernet_custom.config \
  --num_train_steps=50000 \
  --alsologtostderr
  • model_dir: Where training checkpoints and logs will be saved.
  • pipeline_config_path: Path to your modified config file.
5. Monitor Training (Optional)

Track progress with TensorBoard (no Keras callbacks needed):

tensorboard --logdir=./training_output/

You’ll be able to view loss curves, evaluation metrics, and more in your browser.

6. Evaluate and Export the Trained Model

Evaluate the Model

Run evaluation using the same script, pointing to your training checkpoints:

python model_main_tf2.py \
  --model_dir=./training_output/ \
  --pipeline_config_path=./centernet_custom.config \
  --checkpoint_dir=./training_output/ \
  --alsologtostderr

Export for Inference

Export the trained model as a saved model (native TensorFlow format, no Keras wrapper):

python exporter_main_v2.py \
  --input_type=image_tensor \
  --pipeline_config_path=./centernet_custom.config \
  --trained_checkpoint_dir=./training_output/ \
  --output_directory=./exported_centernet/
7. Run Inference with Native TensorFlow

Load the exported model and run detections without any Keras APIs:

import tensorflow as tf
import numpy as np

# Load the saved model (native TensorFlow format)
detect_fn = tf.saved_model.load('./exported_centernet/saved_model')

# Load and preprocess an image
image_path = "test_image.jpg"
image = tf.io.read_file(image_path)
image = tf.image.decode_jpeg(image, channels=3)
input_tensor = tf.convert_to_tensor(image)
input_tensor = input_tensor[tf.newaxis, ...]  # Add batch dimension

# Run inference
detections = detect_fn(input_tensor)

# Process results
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)

# Example: Print detected classes and scores
for i in range(num_detections):
    print(f"Class: {detections['detection_classes'][i]}, Score: {detections['detection_scores'][i]:.2f}")

Quick Notes to Avoid Pitfalls

  • Make sure you’ve properly installed the TensorFlow Object Detection API (compile protobufs, add the repo paths to your PYTHONPATH).
  • If you hit OOM errors, reduce the batch_size in your config or enable mixed precision training by adding use_mixed_precision: true to train_config.

Let me know if you run into specific issues with any step—I’ve debugged my fair share of these setups!

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.29 14:02:29