如何在TensorFlow中实现预测主网络权重的超网络反向传播
解决TensorFlow中超网络权重更新的可微性问题
你遇到的核心问题是tf.assign的特性导致的:它是一个状态更新操作,只会在会话执行时修改变量的值,但不会在计算图中建立从超网络输出到主网络权重的依赖关系。这就导致反向传播时,梯度流到assign这里就中断了,无法传递到超网络的参数上——虽然你能看到权重确实被更新了,但这个更新不在可微的计算路径里。
解决方案:跳过tf.assign,直接构造可微的权重计算路径
正确的做法是不修改主网络的变量状态,而是在计算图中直接用「原始权重 + 超网络预测的增量」作为新的权重,用这个新权重张量来计算主网络的输出。这样整个路径是连续可微的,梯度就能顺利从损失函数传递到超网络的参数。
修改后的代码示例
我们可以调整conv_net函数,让它支持传入超网络生成的权重增量,直接用增量后的权重进行计算:
import numpy as np import tensorflow as tf from tensorflow.contrib.layers import softmax def dense_weight_update_net(inputs, reuse): with tf.variable_scope("weight_net", reuse=reuse): output = tf.layers.conv2d(inputs=inputs, kernel_size=(3, 3), filters=16, strides=(1, 1), activation=tf.nn.leaky_relu, name="conv_layer_0", padding="SAME") output = tf.reduce_mean(output, axis=[0, 1, 2]) output = tf.reshape(output, shape=(1, output.get_shape()[0])) output = tf.layers.dense(output, units=(16*3*3*3)) output = tf.reshape(output, shape=(3, 3, 3, 16)) return output def conv_net(inputs, weight_updates=None, reuse=False): with tf.variable_scope("conv_net", reuse=reuse): # 获取主网络的原始权重变量 conv_kernel = tf.get_variable("conv_layer_0/kernel") conv_bias = tf.get_variable("conv_layer_0/bias") # 如果有超网络的增量,就用原始权重 + 增量作为新权重 if weight_updates is not None: conv_kernel = conv_kernel + weight_updates # 使用更新后的权重执行卷积计算 output = tf.nn.conv2d(inputs=inputs, filters=conv_kernel, strides=(1,1,1,1), padding="SAME") output = tf.nn.bias_add(output, conv_bias) output = tf.nn.leaky_relu(output) output = tf.reduce_mean(output, axis=[1, 2]) output = tf.layers.dense(output, units=2) output = softmax(output) return output # 构建输入和目标 input_x_0 = tf.zeros(shape=(32, 32, 32, 3)) target_y_0 = tf.zeros(shape=(32), dtype=tf.int32) input_x_1 = tf.ones(shape=(32, 32, 32, 3)) target_y_1 = tf.ones(shape=(32), dtype=tf.int32) input_x = tf.concat([input_x_0, input_x_1], axis=0) target_y = tf.concat([target_y_0, target_y_1], axis=0) target_y = tf.one_hot(target_y, 2) # 原始主网络输出(用于对比) output_0 = conv_net(inputs=input_x, reuse=False) crossentropy_loss_0 = tf.losses.softmax_cross_entropy(onehot_labels=target_y, logits=output_0) # 超网络生成权重增量 weight_updates = dense_weight_update_net(inputs=input_x, reuse=False) # 用增量后的权重计算主网络输出(无需assign,直接在计算图中传递依赖) output_1 = conv_net(inputs=input_x, weight_updates=weight_updates, reuse=True) crossentropy_loss_1 = tf.losses.softmax_cross_entropy(onehot_labels=target_y, logits=output_1) # 检查输出差异(验证权重增量生效) check_sum = tf.reduce_sum(tf.abs(output_0 - output_1)) # 定义优化器,只优化超网络参数 weight_net_parameters = tf.get_collection(tf.GraphKeys.TRAINABLE_VARIABLES, scope="weight_net") c_opt = tf.train.AdamOptimizer(beta1=0.9, learning_rate=0.001) train_op = c_opt.minimize(crossentropy_loss_1, var_list=weight_net_parameters) # 训练流程 init = tf.global_variables_initializer() with tf.Session() as sess: sess.run(init) loss_list_0 = [] loss_list_1 = [] for i in range(1000): _, checksum, ce0, ce1 = sess.run([train_op, check_sum, crossentropy_loss_0, crossentropy_loss_1]) loss_list_0.append(ce0) loss_list_1.append(ce1) if i % 100 == 0: print(f"Step {i}: Checksum={checksum:.4f}, Avg Loss0={np.mean(loss_list_0):.4f}, Avg Loss1={np.mean(loss_list_1):.4f}")
关键改动说明
- 移除
tf.assign操作:不再修改主网络变量的状态,而是通过张量运算直接生成更新后的权重。 - 在计算图中建立依赖:超网络的输出
weight_updates直接参与主网络的卷积计算,这样反向传播时梯度可以从crossentropy_loss_1一路传递到超网络的参数。 - 保持主网络变量不变:主网络的原始权重变量不会被修改,每次计算
output_1时都是基于原始权重+超网络增量的临时张量,确保计算图的可微性。
这样调整后,整个系统就可以正常计算超网络的梯度,实现通过超网络优化主网络损失的目标了。
内容的提问来源于stack exchange,提问作者AntreasAntoniou
相关产品推荐
相关产品推荐

