子类初始化设置的stateful属性被super().__init__覆盖问题求助
问题原因与解决方案
这个问题我碰到过,其实核心原因和Keras Layer初始化的逻辑以及Bidirectional的特殊实现有关,我来给你拆解清楚:
为什么你的Dummy类stateful会被覆盖?
你的Dummy继承自Wrapper,而Wrapper又继承自Layer。在Keras 2.1.3的Layer.__init__方法里,有这么一行关键代码:
self.stateful = kwargs.pop('stateful', False)
也就是说:
- 不管你在调用
super().__init__之前怎么设置self.stateful,Layer的初始化过程都会重新给这个属性赋值。 - 如果你调用
super().__init__时没有在kwargs里传入stateful参数,它就会把self.stateful设为默认的False,这就是你的stateful从True变成False的直接原因。
为什么Bidirectional不受影响?
你看到的Bidirectional代码只是一部分——它实际上重写了stateful的@property属性,把这个属性的取值和赋值完全代理给了内部的forward_layer和backward_layer:
@property def stateful(self): return self.forward_layer.stateful @stateful.setter def stateful(self, value): self.forward_layer.stateful = value self.backward_layer.stateful = value
所以即使Layer.__init__把实例变量self.stateful设为False,当你访问bidir.stateful时,调用的是Bidirectional自定义的getter方法,返回的是内部GRU层的stateful值,完全不受Layer初始化的影响。
解决方法
根据你的需求,有三种常用的解决方案:
方案1:在super().__init__之后设置stateful
既然Layer初始化会覆盖这个属性,那我们就在初始化完成后重新赋值,覆盖掉默认值:
from keras.layers import GRU, wrappers class Dummy(wrappers.Wrapper): def __init__(self, layer, **kwargs): super().__init__(layer, **kwargs) self.stateful = True print(self.stateful) # 输出True dummy = Dummy(GRU(64, stateful=True)) print(dummy.stateful) # 输出True
方案2:把stateful=True传入kwargs
直接让Layer的初始化过程使用我们指定的stateful值,避免后续覆盖:
from keras.layers import GRU, wrappers class Dummy(wrappers.Wrapper): def __init__(self, layer, **kwargs): kwargs['stateful'] = True super().__init__(layer, **kwargs) print(self.stateful) # 输出True dummy = Dummy(GRU(64, stateful=True)) print(dummy.stateful) # 输出True
方案3:重写stateful的@property(和Bidirectional逻辑一致)
如果你希望Dummy的stateful和内部包裹的层保持同步,可以像Bidirectional一样使用属性代理:
from keras.layers import GRU, wrappers class Dummy(wrappers.Wrapper): def __init__(self, layer, **kwargs): super().__init__(layer, **kwargs) self.stateful = True # 初始化时设置内部层的stateful @property def stateful(self): return self.layer.stateful @stateful.setter def stateful(self, value): self.layer.stateful = value dummy = Dummy(GRU(64, stateful=True)) print(dummy.stateful) # 输出True
内容的提问来源于stack exchange,提问作者Eli Korvigo
相关产品推荐
相关产品推荐

