计算复杂TF2 Keras模型FLOPS时遇FixedDropout未知层错误
解决EfficientNet计算FLOPS时的FixedDropout报错问题
报错原因
EfficientNet模型内置了FixedDropout这个Keras自定义层,TensorFlow在序列化模型或转换计算图时,无法自动识别未注册的自定义层,因此抛出Unknown layer: FixedDropout错误。
针对第一种FLOPS计算函数(直接传入model对象)
修改步骤:
- 先导入
FixedDropout层:
from tensorflow.keras.applications.efficientnet import FixedDropout
- 在调用
get_flops函数前,将该自定义层注册到Keras的自定义对象库中:
tf.keras.utils.get_custom_objects()['FixedDropout'] = FixedDropout
- 之后正常调用原
get_flops函数即可。
如果想把注册逻辑整合到函数内部,修改后的函数如下:
from tensorflow.keras.applications.efficientnet import FixedDropout from tensorflow.python.framework.convert_to_constants import convert_variables_to_constants_v2_as_graph import tensorflow as tf def get_flops(model, batch_size=None): # 注册自定义层 tf.keras.utils.get_custom_objects()['FixedDropout'] = FixedDropout if batch_size is None: batch_size = 1 real_model = tf.function(model).get_concrete_function(tf.TensorSpec([batch_size] + model.inputs[0].shape[1:], model.inputs[0].dtype)) frozen_func, graph_def = convert_variables_to_constants_v2_as_graph(real_model) run_meta = tf.compat.v1.RunMetadata() opts = tf.compat.v1.profiler.ProfileOptionBuilder.float_operation() flops = tf.compat.v1.profiler.profile(graph=frozen_func.graph, run_meta=run_meta, cmd='op', options=opts) return flops.total_float_ops
针对第二种FLOPS计算函数(从h5文件加载模型)
核心是在load_model时显式传入custom_objects参数,注册FixedDropout层。修改后的函数如下:
from tensorflow.keras.applications.efficientnet import FixedDropout import tensorflow as tf def get_flops(model_h5_path): session = tf.compat.v1.Session() graph = tf.compat.v1.get_default_graph() with graph.as_default(): with session.as_default(): # 加载模型时传入自定义层 model = tf.keras.models.load_model(model_h5_path, custom_objects={'FixedDropout': FixedDropout}) run_meta = tf.compat.v1.RunMetadata() opts = tf.compat.v1.profiler.ProfileOptionBuilder.float_operation() flops = tf.compat.v1.profiler.profile(graph=graph, run_meta=run_meta, cmd='op', options=opts) return flops.total_float_ops
通用注意事项
- 如果使用的是EfficientNetV2,对应的自定义 dropout 层是
StochasticDepth,需要从tf.keras.applications.efficientnet_v2导入并注册,方法同上。 - 确保TensorFlow版本与EfficientNet的版本兼容,避免因版本差异导致的层定义不一致问题。
内容的提问来源于stack exchange,提问作者Imperial_J
相关产品推荐
相关产品推荐

