如何将(Fused)BatchNorm转为卷积/加法以在TensorFlow Lite运行自定义MobileNet
解决自定义MobileNet的(Fused)BatchNorm适配TensorFlow Lite的问题
我之前折腾自定义MobileNet转TFLite时也碰到过几乎一模一样的问题,核心就是早期TOCO对FusedBatchNorm的支持拉胯,和教程里拆分的BN结构不兼容。下面给你几个亲测可行的解决路径:
优先用官方新转换工具替代TOCO
TOCO早就被整合到TensorFlow的tf.lite.TFLiteConverter里了,这个工具对FusedBatchNorm的处理要成熟得多,基本能自动搞定兼容问题。步骤很简单:
- 加载你的模型(不管是SavedModel还是冻结pb):
import tensorflow as tf # 从SavedModel加载(推荐,训练后直接存成SavedModel更方便) converter = tf.lite.TFLiteConverter.from_saved_model("你的SavedModel路径") # 或者从冻结图加载 # converter = tf.lite.TFLiteConverter.from_frozen_graph( # graph_def_file="frozen_model.pb", # input_arrays=["输入节点名称"], # output_arrays=["输出节点名称"] # ) - 开启默认优化(可选,但能兼顾兼容性和移动端性能):
converter.optimizations = [tf.lite.Optimize.DEFAULT] - 转换并保存TFLite模型:
tflite_model = converter.convert() with open("custom_mobilenet.tflite", "wb") as f: f.write(tflite_model)
用这个工具转换的话,基本不需要手动改BN结构,它会自动处理FusedBatchNorm到TFLite兼容算子的转换。
训练时直接用非融合BN结构
如果你一定要和教程里的模型结构对齐,那可以在训练阶段就把BatchNorm设置成非融合模式:
- 在定义MobileNet的BatchNorm层时,给
BatchNormalization加上fused=False参数:from tensorflow.keras.layers import BatchNormalization # 替换你模型里所有的BatchNorm层 bn_layer = BatchNormalization(fused=False, epsilon=1e-3, momentum=0.999, ...) - 至于你之前碰到的
SquaredDifference算子不支持问题,大概率是导出模型时带了训练相关的节点(比如损失计算节点),导出冻结图时一定要明确指定只保留推理需要的输入输出节点,把无关节点过滤掉。比如用tf.compat.v1.graph_util.convert_variables_to_constants时,准确设置output_node_names参数。
手动修改冻结图拆分FusedBatchNorm(不推荐,除非不想重训)
如果已经训好模型不想再来一遍,那可以手动修改pb文件的图结构,把FusedBatchNorm拆成TFLite支持的基础算子组合:
- 先加载冻结的图定义:
import tensorflow as tf from tensorflow.core.framework import graph_pb2 with open("frozen_model.pb", "rb") as f: graph_def = graph_pb2.GraphDef() graph_def.ParseFromString(f.read()) - 遍历图里的所有节点,找到
FusedBatchNorm节点,把它替换成Mul(缩放)+Add(偏移)的组合——其实FusedBatchNorm的计算逻辑就是:output = gamma * (input - moving_mean) / sqrt(moving_variance + epsilon) + beta,用几个基础算子就能拼出来,这些算子都是TFLite支持的。 - 修改完图结构后,重新保存pb,再用TFLite Converter转换就行。
最后提个小建议
尽量用TensorFlow最新的稳定版,旧版本对TFLite的算子支持有很多坑,升级后很多奇怪的错误会直接消失。另外转换前可以用tf.lite.experimental.Analyzer检查模型里的算子兼容性:
tf.lite.experimental.Analyzer.analyze(model_content=tflite_model)
内容的提问来源于stack exchange,提问作者rhavard
相关产品推荐
相关产品推荐

