如何在自定义Keras层中定义不被跟踪的非持久化变量
Keras自定义层排除变量追踪解决方案
核心原理
Keras层实例默认会自动追踪所有直接赋值给实例属性的tf.Variable对象,无论是否可训练,都会被纳入层的权重列表、随模型权重一同持久化。要避免该行为,只需打破Keras的直接追踪逻辑即可。
实现方法
方法1:使用官方追踪排除接口(推荐,兼容性好)
调用Keras层内置的_tracker.ignore()方法标记不需要追踪的变量,修改后的代码如下:
class Codebook(layers.Layer): def __init__(self, num_codes, code_reset_limit = None, **kwargs): super().__init__(**kwargs) self.num_codes = num_codes self.code_reset_limit = code_reset_limit if self.code_reset_limit: self.code_counter = tf.Variable(tf.zeros(num_codes, dtype = tf.int32), trainable = False) # 标记该变量不被Keras追踪 self._tracker.ignore(self.code_counter) def build(self, input_shape): self.codes = self.add_weight(name = 'codes', shape = (self.num_codes, input_shape[-1]), initializer = 'random_uniform', trainable = True) super().build(input_shape)
使用该方法可以保持原有调用self.code_counter的逻辑不变,无需修改其他业务代码。
方法2:将变量存入普通Python容器
Keras不会追踪普通Python容器(列表、字典等)内部的元素,适合有多个辅助变量的场景,修改后的代码如下:
class Codebook(layers.Layer): def __init__(self, num_codes, code_reset_limit = None, **kwargs): super().__init__(**kwargs) self.num_codes = num_codes self.code_reset_limit = code_reset_limit # 用普通列表存储辅助变量 self._aux_vars = [] if self.code_reset_limit: code_counter = tf.Variable(tf.zeros(num_codes, dtype = tf.int32), trainable = False) self._aux_vars.append(code_counter) def build(self, input_shape): self.codes = self.add_weight(name = 'codes', shape = (self.num_codes, input_shape[-1]), initializer = 'random_uniform', trainable = True) super().build(input_shape)
后续需要使用计数器时,通过self._aux_vars[0]调用即可。
注意事项
- 未被追踪的变量不会随模型权重保存,每次训练重启时计数器会从零开始计数,符合推理阶段不需要码本重置逻辑的需求
- 若使用TensorFlow 2.4以下的旧版本,
_tracker接口可能不存在,直接使用方法2即可保证兼容性
内容的提问来源于stack exchange,提问作者chasep255
相关产品推荐
相关产品推荐

