如何配置TFLite Model Maker实现EfficientDet-Lite3离线训练?
离线配置TFLite Model Maker的EfficientDet-Lite3目标检测训练环境
问题背景
使用TFLite Model Maker的EfficientDet-Lite3进行目标检测时,在线环境下基础训练流程可正常运行,但离线机器上因model_spec.get('efficientdet_lite3')需要联网获取预训练模型,出现urlopen错误。已下载模型文件,但尝试用EfficientDetLite3Spec指向本地SavedModel时触发以下错误:
ERROR:absl:hub.KerasLayer is trainable but has zero trainable weights. ValueError: Could not find matching concrete function to call loaded from the SavedModel. Got: Positional arguments (1 total): * <tf.Tensor 'imgs:0' shape=(None, 512, 512, 3) dtype=float32> Keyword arguments: {} Expected these arguments to match one of the following 1 option(s):
原训练代码:
import numpy as np import os from tflite_model_maker.config import ExportFormat, QuantizationConfig from tflite_model_maker import model_spec from tflite_model_maker import object_detector from tflite_support import metadata import tensorflow as tf assert tf.__version__.startswith('2') tf.get_logger().setLevel('ERROR') from absl import logging logging.set_verbosity(logging.ERROR) train_data = object_detector.DataLoader.from_pascal_voc( 'my_data/train', 'my_data/train', ['obj'] ) val_data = object_detector.DataLoader.from_pascal_voc( 'my_data/validate', 'my_data/validate', ['obj'] ) spec = model_spec.get('efficientdet_lite3') model = object_detector.create(train_data, model_spec=spec, batch_size=4, train_whole_model=True, epochs=200, validation_data=val_data)
解决方案
方法1:迁移联网机器的模型缓存(推荐)
- 在联网机器上运行一次
spec = model_spec.get('efficientdet_lite3'),触发预训练模型自动下载,完成后关闭程序 - 找到模型缓存目录:默认路径为
~/.cache/tensorflow/hub/,其中会有一个以模型TFHub地址命名的文件夹(如https://tfhub.dev/tensorflow/efficientdet/lite3/feature-vector/1) - 将该文件夹完整复制到离线机器的相同缓存路径,或设置环境变量
TFHUB_CACHE_DIR指定自定义缓存路径,再把模型文件夹放入该路径 - 离线机器上运行原训练代码,即可直接加载本地缓存的模型
方法2:自定义EfficientDetLite3Spec加载本地模型
若已拥有符合要求的预训练SavedModel,需确保模型输入输出与TFLite Model Maker的训练逻辑匹配,修改代码如下:
import numpy as np import os from tflite_model_maker.config import ExportFormat, QuantizationConfig from tflite_model_maker import object_detector from tflite_model_maker.object_detector import EfficientDetLite3Spec # 导入Spec类 from tflite_support import metadata import tensorflow as tf assert tf.__version__.startswith('2') tf.get_logger().setLevel('ERROR') from absl import logging logging.set_verbosity(logging.ERROR) train_data = object_detector.DataLoader.from_pascal_voc( 'my_data/train', 'my_data/train', ['obj'] ) val_data = object_detector.DataLoader.from_pascal_voc( 'my_data/validate', 'my_data/validate', ['obj'] ) # 创建自定义spec,指向本地SavedModel路径 spec = EfficientDetLite3Spec( uri='./path/to/your/local/efficientdet_lite3_saved_model', # 替换为本地模型路径 input_image_shape=(512, 512), # 匹配训练数据的输入尺寸,需与模型要求一致 model_name='efficientdet-lite3' ) model = object_detector.create(train_data, model_spec=spec, batch_size=4, train_whole_model=True, epochs=200, validation_data=val_data)
注意:需确保本地SavedModel是训练专用版本(包含可训练变量和正确的输入签名),避免使用仅用于推理的导出模型。
额外注意事项
- 确保离线机器的TensorFlow、TFLite Model Maker版本与联网机器完全一致,避免版本兼容问题
- 若使用自定义缓存路径,需在运行代码前设置环境变量:
export TFHUB_CACHE_DIR="/your/custom/cache/path"(Linux/macOS)或set TFHUB_CACHE_DIR="C:\your\custom\cache\path"(Windows)
内容的提问来源于stack exchange,提问作者Dodo
相关产品推荐
相关产品推荐

