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

如何配置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:迁移联网机器的模型缓存(推荐)

  1. 在联网机器上运行一次spec = model_spec.get('efficientdet_lite3'),触发预训练模型自动下载,完成后关闭程序
  2. 找到模型缓存目录:默认路径为~/.cache/tensorflow/hub/,其中会有一个以模型TFHub地址命名的文件夹(如https://tfhub.dev/tensorflow/efficientdet/lite3/feature-vector/1)
  3. 将该文件夹完整复制到离线机器的相同缓存路径,或设置环境变量TFHUB_CACHE_DIR指定自定义缓存路径,再把模型文件夹放入该路径
  4. 离线机器上运行原训练代码,即可直接加载本地缓存的模型

方法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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 05:35:09