如何基于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:
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
- For Pascal VOC-style datasets:
- If your dataset uses a custom format, modify one of these scripts to parse your annotations and images into TFRecords.
- Create a
label_map.pbtxtfile that maps class IDs to names, like this:item { id: 1 name: 'your_class_1' } item { id: 2 name: 'your_class_2' }
Grab the pre-trained checkpoint from TensorFlow Model Zoo:
- Download the Centernet_1024x1024 package (look for the
.tar.gzfile containing the checkpoint) - Extract it to a directory (e.g.,
./pretrained_centernet/)—you’ll need theckpt-xxxxfiles inside for fine-tuning.
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_pathto point to your training TFRecord file. - Update
label_map_pathto yourlabel_map.pbtxtfile.
- Update
eval_input_reader: Do the same as above for your validation TFRecord.train_config:- Set
fine_tune_checkpointto 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, andnum_stepsbased on your hardware and dataset size.
- Set
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.
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.
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/
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_sizein your config or enable mixed precision training by addinguse_mixed_precision: truetotrain_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

