使用TensorFlow EfficientNetB0创建特征提取模型时遇类型错误
问题:创建EfficientNetB0特征提取模型时的 dtype 不匹配错误
错误信息
--------------------------------------------------------------------------- TypeError Traceback (most recent call last) <ipython-input-23-a777a6ba699b> in <module> ----> 1 model_00 = create_feature_extraction_model() 2 model_00.summary() <ipython-input-22-b1a80357611c> in create_feature_extraction_model() 1 def create_feature_extraction_model(): ----> 2 base_model = tf.keras.applications.EfficientNetB0(include_top=False) 3 4 base_model.trainable = False 5 ~/anaconda3/envs/ml/lib/python3.8/site-packages/keras/applications/efficientnet.py in EfficientNetB0(include_top, weights, input_tensor, input_shape, pooling, classes, classifier_activation, **kwargs) 546 classifier_activation='softmax', 547 **kwargs): --> 548 return EfficientNet( 549 1.0, 550 1.0, ~/anaconda3/envs/ml/lib/python3.8/site-packages/keras/applications/efficientnet.py in EfficientNet(width_coefficient, depth_coefficient, default_size, dropout_rate, drop_connect_rate, depth_divisor, activation, blocks_args, model_name, include_top, weights, input_tensor, input_shape, pooling, classes, classifier_activation) 332 # original implementation. 333 # See https://github.com/tensorflow/tensorflow/issues/49930 for more details --> 334 x = x / tf.math.sqrt(IMAGENET_STDDEV_RGB) 335 336 x = layers.ZeroPadding2D( ~/anaconda3/envs/ml/lib/python3.8/site-packages/tensorflow/python/util/traceback_utils.py in error_handler(*args, **kwargs) 151 except Exception as e: 152 filtered_tb = _process_traceback_frames(e.__traceback__) --> 153 raise e.with_traceback(filtered_tb) from None 154 finally: 155 del filtered_tb ~/anaconda3/envs/ml/lib/python3.8/site-packages/keras/layers/core/tf_op_layer.py in handle(self, op, args, kwargs) 105 isinstance(x, keras_tensor.KerasTensor) 106 for x in tf.nest.flatten([args, kwargs])): --> 107 return TFOpLambda(op)(*args, **kwargs) 108 else: 109 return self.NOT_SUPPORTED ~/anaconda3/envs/ml/lib/python3.8/site-packages/keras/utils/traceback_utils.py in error_handler(*args, **kwargs) 65 except Exception as e: # pylint: disable=broad-except 66 filtered_tb = _process_traceback_frames(e.__traceback__) ---> 67 raise e.with_traceback(filtered_tb) from None 68 finally: 69 del filtered_tb TypeError: Exception encountered when calling layer "tf.math.truediv" (type TFOpLambda). `x` and `y` must have the same dtype, got tf.float16 != tf.float32. Call arguments received by layer "tf.math.truediv" (type TFOpLambda): • x=tf.Tensor(shape=(None, None, None, 3), dtype=float16) • y=tf.Tensor(shape=(3,), dtype=float32) • name=None
环境与代码细节
- 数据集:tensorflow_datasets的
food101,图像形状(224,224,3),dtype=tf.float32 - 训练配置:启用
mixed_float16混合精度训练 - GPU规格:
+-----------------------------------------------------------------------------+ | NVIDIA-SMI 520.61.05 Driver Version: 520.61.05 CUDA Version: 11.8 | |-------------------------------+----------------------+----------------------+ | GPU Name Persistence-M| Bus-Id Disp.A | Volatile Uncorr. ECC | | Fan Temp Perf Pwr:Usage/Cap| Memory-Usage | GPU-Util Compute M. | | | | MIG M. | |===============================+======================+======================| | 0 NVIDIA GeForce ... On | 00000000:01:00.0 Off | N/A | | N/A 56C P0 13W / N/A | 7MiB / 4096MiB | 0% Default | | | | N/A | +-------------------------------+----------------------+----------------------+ +-----------------------------------------------------------------------------+ | Processes: | | GPU GI CI PID Type Process name GPU Memory | | ID ID Usage | |=============================================================================| | 0 N/A N/A 6850 G /usr/lib/xorg/Xorg 4MiB | +-----------------------------------------------------------------------------+
- 模型代码:
def create_feature_extraction_model(): base_model = tf.keras.applications.EfficientNetB0(include_top=False) base_model.trainable = False inputs = tf.keras.layers.Input(shape=IMG_SHAPE, name='input_layer') x = base_model(inputs, training=False) x = tf.keras.layers.GlobalAvgPool2D()(x) x = tf.keras.layers.Dense(units=num_classes)(x) outputs = tf.keras.layers.Softmax(dtype=tf.float32, name='output_layer')(x) model = tf.keras.models.Model(inputs, outputs, name='food101_model') model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy']) return model
解决方案
错误根源是混合精度训练模式下,EfficientNet内部的归一化常量IMAGENET_STDDEV_RGB为float32类型,但输入被自动转为float16,导致 dtype 不匹配。以下是三种可行的解决方法:
方法1:显式指定基础模型的 dtype
创建EfficientNetB0时直接指定dtype=tf.float16,让模型内部常量自动匹配 dtype:
base_model = tf.keras.applications.EfficientNetB0(include_top=False, dtype=tf.float16)
方法2:临时关闭混合精度构建基础模型
构建基础模型前临时切换到float32策略,完成后恢复原混合精度策略:
def create_feature_extraction_model(): # 临时关闭混合精度 from tensorflow.keras import mixed_precision policy = mixed_precision.global_policy() mixed_precision.set_global_policy('float32') base_model = tf.keras.applications.EfficientNetB0(include_top=False) # 恢复混合精度策略 mixed_precision.set_global_policy(policy) base_model.trainable = False inputs = tf.keras.layers.Input(shape=IMG_SHAPE, name='input_layer') x = base_model(inputs, training=False) x = tf.keras.layers.GlobalAvgPool2D()(x) x = tf.keras.layers.Dense(units=num_classes)(x) outputs = tf.keras.layers.Softmax(dtype=tf.float32, name='output_layer')(x) model = tf.keras.models.Model(inputs, outputs, name='food101_model') model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy']) return model
方法3:手动统一输入 dtype
在输入层后添加类型转换层,确保输入与模型内部常量 dtype 一致:
inputs = tf.keras.layers.Input(shape=IMG_SHAPE, name='input_layer') # 转换为float32,匹配模型内部常量类型 x = tf.keras.layers.Lambda(lambda x: tf.cast(x, tf.float32))(inputs) x = base_model(x, training=False)
内容的提问来源于stack exchange,提问作者Iniyan Kanmani
相关产品推荐
相关产品推荐

