TensorFlow2.3自定义训练专属图像增强层报错求助
我明白你遇到的问题了——在函数式API里用自定义的图像增强层时,因为training=None参数导致了"None values not supported"的错误,而Sequential却能正常运行。这其实是因为Keras在不同模型结构下处理training参数的方式不一样。
问题根源
你在call方法里定义了training=None参数,但在函数式API构建模型的过程中,这个参数并没有被Keras自动注入正确的训练/推断状态(True/False),而是保持了None值,传给tf.cond后就触发了错误。而Sequential会自动帮你处理这个参数的传递,所以没出问题。
修复方案
有两种简单的方式解决这个问题,推荐第一种,更简洁且符合Keras的设计规范:
方案1:利用Layer内置的self.training属性
Keras的Layer类本身就自带了training属性,它会自动根据模型当前的运行模式(训练/推断)切换True/False状态。你只需要把call方法里的training参数去掉,改用self.training即可:
class RandomLight(layers.Layer): def __init__(self, factor=0.2): super(RandomLight,self).__init__() self.factor = factor def call(self, input): # 直接使用self.training判断当前模式 if self.training: return tf.clip_by_value(tf.image.random_brightness(input, self.factor), 0, 1) else: return input
或者保持tf.cond的写法也可以:
def call(self, input): return tf.cond( self.training, lambda: tf.clip_by_value(tf.image.random_brightness(input, self.factor), 0, 1), lambda: input )
这样修改后,当你用函数式API构建模型时,Keras会自动管理self.training的值:
- 调用
model.fit()时,self.training为True,执行图像增强 - 调用
model.predict()或model.evaluate()时,self.training为False,直接返回原图
方案2:显式传递training参数到层
如果你想手动控制training状态,可以在函数式API构建时显式传递这个参数,同时保留call方法里的training参数:
首先修改层的call方法:
class RandomLight(layers.Layer): def __init__(self, factor=0.2): super(RandomLight,self).__init__() self.factor = factor def call(self, input, training=None): # 这里可以给training设置默认的判断逻辑 training = training if training is not None else self.training return tf.cond( training, lambda: tf.clip_by_value(tf.image.random_brightness(input, self.factor), 0, 1), lambda: input )
然后在函数式API中构建模型时,传递training参数:
from tensorflow.keras import layers, Model input_layer = layers.Input(shape=(224, 224, 3)) # 训练时传递training=True,推断时传递training=False x = RandomLight()(input_layer, training=True) # 后续层... output_layer = layers.Dense(10, activation='softmax')(x) model = Model(inputs=input_layer, outputs=output_layer)
如果想让training参数自动跟随Keras的全局学习状态,可以用K.learning_phase():
from tensorflow.keras import backend as K x = RandomLight()(input_layer, training=K.learning_phase())
验证
修改后,你可以用函数式API正常构建模型,训练时会自动执行图像增强,推断时则不会,而且不会再出现None values not supported的错误。
内容的提问来源于stack exchange,提问作者dFWXCff

