TensorFlow变量保存疑问:赋值变量及ResNet三元权重存Checkpoint异常
我来帮你解决这两个TensorFlow的问题,都是变量保存相关的常见场景,咱们逐个拆解:
在TensorFlow里,要保存通过其他变量赋值得到的变量,核心是确保这个目标变量能被Checkpoint正确追踪。这里分两种常见情况说明:
如果目标变量是你显式创建的
tf.Variable:
比如你先定义了变量var_a,再创建var_b并通过assign把var_a的值赋给它。这种情况下,只要把var_b加入到Checkpoint的追踪对象里就行。举个例子:# 初始化变量并完成赋值 var_a = tf.Variable(tf.random.normal([3, 3])) var_b = tf.Variable(tf.zeros_like(var_a)) var_b.assign(var_a) # 保存变量到Checkpoint checkpoint = tf.train.Checkpoint(var_a=var_a, var_b=var_b) checkpoint.save('./my_checkpoint_dir')后续加载Checkpoint时,
var_b就能恢复到赋值后的状态。如果目标变量是自定义层/模型中的赋值变量:
要是你在自定义层里基于层的权重变量赋值得到新变量,记得把这个新变量设为层的类属性(比如self.target_var),这样模型的Checkpoint会自动追踪到层的所有属性变量,不用手动指定。
根据你描述的情况,Tenary_W能输出-1、0、+1但存不进Checkpoint,大概率是因为它是临时计算的张量,而不是可被Checkpoint追踪的tf.Variable——毕竟Checkpoint只认tf.Variable对象,临时张量是不会被保存的。这里给你三个可行的解决方案:
方案1:把三元化权重存为自定义层的可追踪变量
修改你的Tenary_Conv2D层,把三元化后的权重作为层的属性变量,而不是每次前向传播时临时生成。示例代码大概是这样:
class Tenary_Conv2D(tf.keras.layers.Conv2D): def build(self, input_shape): super().build(input_shape) # 创建一个非训练变量来存储三元化后的权重 self.ternary_w = tf.Variable(tf.zeros_like(self.kernel), trainable=False) def call(self, inputs): # 执行三元化逻辑,得到临时张量 threshold = tf.reduce_mean(tf.abs(self.kernel)) ternary_w_tensor = tf.where( self.kernel > threshold, 1.0, tf.where(self.kernel < -threshold, -1.0, 0.0) ) # 将三元化结果赋值给层的属性变量 self.ternary_w.assign(ternary_w_tensor) # 使用三元化后的权重做卷积运算 return tf.nn.conv2d( inputs, self.ternary_w, strides=self.strides, padding=self.padding, data_format=self.data_format )
这样self.ternary_w是层的可追踪变量,保存模型Checkpoint时会自动被纳入保存范围。
方案2:保存前把三元化权重赋值回原始权重变量
如果你不想额外创建变量,也可以在保存Checkpoint前,把三元化后的权重覆盖到卷积层的原始kernel变量上,再执行保存:
# 遍历模型中的所有Tenary_Conv2D层 for layer in model.layers: if isinstance(layer, Tenary_Conv2D): # 先获取三元化后的权重张量(需要在层里加一个方法返回这个张量) ternary_w_tensor = layer.get_ternary_weight() # 你需要自己实现这个方法 # 把三元化结果赋值给原始kernel变量 layer.kernel.assign(ternary_w_tensor) # 保存模型Checkpoint checkpoint = tf.train.Checkpoint(model=model) checkpoint.save('./resnet_ternary_checkpoint')
这个方法的好处是不用新增变量,但要注意:赋值会覆盖原始权重,要是你还需要保留原始权重,记得提前备份。
方案3:自定义Checkpoint的保存变量列表
你可以手动收集所有需要保存的变量(包括三元化权重变量),然后传给Checkpoint:
# 收集所有Tenary_Conv2D层的三元化权重变量 ternary_vars = [ layer.ternary_w for layer in model.layers if isinstance(layer, Tenary_Conv2D) ] # 合并模型的可训练变量和三元化变量 all_save_vars = model.trainable_variables + ternary_vars # 创建Checkpoint并保存 checkpoint = tf.train.Checkpoint(variables=all_save_vars) checkpoint.save('./custom_ternary_checkpoint')
另外,你提到权重直方图能看到-1、0、+1,说明三元化逻辑是没问题的,只是没把结果存成可保存的变量,按照上面的方法调整后应该就能正常保存了。
内容的提问来源于stack exchange,提问作者chao wang

