Keras子类化RandomRotation类时出现未知关键字参数报错
解决Keras子类化RandomRotation时的参数错误问题
问题根源
你遇到的TypeError是因为value_range和data_format并非RandomRotation类初始化方法(__init__)的合法参数,这两个参数属于调用层处理图像时(call方法)的输入参数,而非实例化时的配置参数。当你把它们放在**kwargs中传递给父类的__init__时,父类无法识别这些参数,因此抛出错误。
解决方法
方法1:修正实例化代码,移除非法初始化参数
直接移除value_range和data_format,仅保留RandomRotation支持的初始化参数:
ir_layer = ImageRotator( factor=(0,1), fill_mode="constant", interpolation="bilinear", seed=1, fill_value=0.0, )
如果需要指定value_range和data_format,在调用层处理图像时传入:
# 假设img是输入图像张量 rotated_img = ir_layer(img, value_range=(0, 255), data_format=None)
方法2:修改子类,支持初始化时传入这些参数
如果希望子类在实例化时就能接收value_range和data_format,可以显式在子类__init__中声明这些参数,存储为实例属性后,在call方法中传递给父类:
from keras.layers import RandomRotation class ImageRotator(RandomRotation): def __init__(self, factor, value_range=None, data_format=None, **kwargs): super().__init__(factor=factor, **kwargs) self.value_range = value_range self.data_format = data_format def call(self, inputs, training=True): return super().call( inputs, training=training, value_range=self.value_range, data_format=self.data_format )
此时实例化代码即可正常传入所有参数:
ir_layer = ImageRotator( factor=(0,1), fill_mode="constant", interpolation="bilinear", seed=1, fill_value=0.0, value_range=(0, 255), data_format=None, )
内容的提问来源于stack exchange,提问作者statsman92
相关产品推荐
相关产品推荐

