如何将TensorFlow训练的SSD MobileNetV2模型转换为CoreML
报错根因
tf.saved_model.load() 加载TensorFlow目标检测API导出的SavedModel时,返回的是_UserObject类型的包装对象,不在coremltools 5.2支持的输入格式范围内(支持格式为SavedModel路径、concrete function列表、tf.keras.Model对象、h5权重文件),直接传入该对象必然触发类型不匹配报错。同时目标检测API导出的模型默认没有封装为Keras模型结构,无法被coremltools自动识别解析。
可复现的正确转换流程
- 导出适配转换的SavedModel
调用官方exporter_main_v2.py导出checkpoint时,指定输入类型为浮点图像张量,避免后续输入维度解析异常,参考命令:
不要使用python exporter_main_v2.py \ --trained_checkpoint_dir=<你的本地checkpoint存储路径> \ --pipeline_config_path=<训练用pipeline.config文件路径> \ --output_directory=./exported_det_model \ --input_type=float_image_tensorinput_type=image_tensor参数导出,该模式导出的输入为uint8编码的图像张量,coremltools解析时容易出现维度匹配错误。 - 提取可被coremltools识别的推理函数
加载SavedModel后不要直接传加载返回的根对象,提取默认服务签名对应的concrete function作为转换输入:
SSD MobileNetV2默认输入形状为import tensorflow as tf import coremltools as ct # 加载导出的SavedModel loaded = tf.saved_model.load("./exported_det_model/saved_model") # 提取推理用concrete function,目标检测API导出的默认推理签名key固定为serving_default infer_func = loaded.signatures["serving_default"] # 可选校验:打印输入输出结构,确认和训练配置对齐 print("输入结构:", infer_func.structured_input_signature) print("输出结构:", infer_func.structured_outputs)(1, 宽, 高, 3),宽高值和你训练配置中设定的输入尺寸一致,常见为320*320。 - 执行CoreML转换
转换时传入提取到的concrete function,同时明确输入类型,避免自动解析失败:# 定义输入张量类型,scale参数对应0-255到0-1的归一化逻辑,和模型训练时的预处理对齐 input_tensor = ct.ImageType( name="input_tensor", shape=(1, 320, 320, 3), scale=1/255.0, bias=[0, 0, 0] ) mlmodel = ct.convert( infer_func, source="tensorflow", inputs=[input_tensor], convert_to="mlprogram", compute_units=ct.ComputeUnit.ALL ) - (可选)优化输出命名并保存
TF目标检测API导出的模型默认输出四个张量,可重命名后保存方便端侧调用:# 重命名输出 mlmodel.output_names["detection_boxes"] = "detection_boxes" mlmodel.output_names["detection_classes"] = "detection_classes" mlmodel.output_names["detection_scores"] = "detection_scores" mlmodel.output_names["num_detections"] = "num_detections" # 保存为mlpackage格式 mlmodel.save("./ssd_mobilenetv2_detector.mlpackage")
注意事项
- 不要尝试用
tf.keras.models.load_model加载目标检测API导出的SavedModel,这类模型不是通过Keras接口保存的,强行加载会触发结构解析错误。 - 若转换过程中遇到算子不支持报错,可升级coremltools到5.2最新补丁版本,SSD MobileNetV2用到的卷积、锚框解码、NMS相关算子在5.2版本已完成原生适配。
- 转换前确认导出的SavedModel可以在TensorFlow环境下正常跑通推理,避免因checkpoint导出损坏导致转换失败。
内容的提问来源于stack exchange,提问作者manozd
相关产品推荐
相关产品推荐

