TensorFlow不同variable_scope下变量共享问题:Actor-Critic网络层共享实现
我之前在做Actor-Critic架构的时候也踩过这个变量共享的坑,尤其是用TensorFlow 1.x的variable_scope时很容易搞混。针对你需要的V与Actor共享部分网络层、V_target是V的精确副本的需求,这里给你一套靠谱的实现方案,分两种主流框架版本来说:
一、TensorFlow 1.x 实现方案
1. 先抽离共享的基础网络模块
把要共享的层单独写成一个函数,通过reuse参数控制变量的创建与复用——第一次调用(给V网络)时创建变量,后续Actor调用时复用这些变量:
import tensorflow as tf def shared_backbone(inputs, reuse=False): # 这个scope是全局的,不要嵌套在其他隔离scope里 with tf.variable_scope("shared_layers", reuse=reuse): hidden1 = tf.layers.dense(inputs, 64, activation=tf.nn.relu, name="h1") hidden2 = tf.layers.dense(hidden1, 64, activation=tf.nn.relu, name="h2") return hidden2
2. 构建V主网络与Actor网络
V网络先调用共享模块创建变量,Actor网络调用时开启reuse=True,就能复用共享层的变量了:
# 定义输入占位符 state_dim = 8 # 替换成你的状态维度 action_dim = 4 # 替换成你的动作维度 state_input = tf.placeholder(tf.float32, shape=[None, state_dim], name="state") # 构建V主网络 shared_features_v = shared_backbone(state_input, reuse=False) with tf.variable_scope("value_main"): v_pred = tf.layers.dense(shared_features_v, 1, name="value_out") # 构建Actor网络,复用共享层变量 shared_features_actor = shared_backbone(state_input, reuse=True) with tf.variable_scope("actor"): action_logits = tf.layers.dense(shared_features_actor, action_dim, name="action_logits")
这里的关键是:共享层的scope不能被嵌套在V或Actor的专属scope里,否则会生成value_main/shared_layers和actor/shared_layers两个独立的变量集合,完全无法共享——这应该就是你之前失败的原因!
3. 构建V_target目标网络
V_target是V的精确副本,我们不需要重新定义网络结构(避免重复创建变量),而是通过变量赋值来同步参数:
# 获取V主网络的所有可训练变量(共享层+V专属层) v_main_vars = tf.get_collection( tf.GraphKeys.TRAINABLE_VARIABLES, scope="shared_layers|value_main" ) # 构建V_target的结构,注意scope要和V主网络对应 def build_value_target(inputs): with tf.variable_scope("value_target"): h1 = tf.layers.dense(inputs, 64, activation=tf.nn.relu, name="h1") h2 = tf.layers.dense(h1, 64, activation=tf.nn.relu, name="h2") v_target_pred = tf.layers.dense(h2, 1, name="value_out") return v_target_pred v_target_pred = build_value_target(state_input) v_target_vars = tf.get_collection( tf.GraphKeys.TRAINABLE_VARIABLES, scope="value_target" ) # 创建同步操作:把V主网络的参数赋值给V_target update_v_target = [tf.assign(t_var, m_var) for t_var, m_var in zip(v_target_vars, v_main_vars)] # 初始化时先同步一次,保证初始参数一致 init_v_target = tf.group(*update_v_target)
训练时,每隔固定步数(比如1000步)调用一次init_v_target的会话运行操作,就能保持V_target和V主网络的参数同步。
二、TensorFlow 2.x/Keras 实现方案
TF2.x用Keras的话会简单很多,直接把共享层定义成对象,重复使用即可,完全不用管variable_scope的问题:
import tensorflow as tf # 定义共享层对象 shared_dense1 = tf.keras.layers.Dense(64, activation='relu') shared_dense2 = tf.keras.layers.Dense(64, activation='relu') # 构建V主网络 state_input = tf.keras.Input(shape=(state_dim,)) x = shared_dense1(state_input) x = shared_dense2(x) v_out = tf.keras.layers.Dense(1)(x) v_main_net = tf.keras.Model(inputs=state_input, outputs=v_out) # 构建Actor网络,直接复用共享层对象 x = shared_dense1(state_input) x = shared_dense2(x) action_out = tf.keras.layers.Dense(action_dim)(x) actor_net = tf.keras.Model(inputs=state_input, outputs=action_out) # 构建V_target:直接克隆V主网络并同步参数 v_target_net = tf.keras.models.clone_model(v_main_net) v_target_net.set_weights(v_main_net.get_weights())
训练时,同步V_target只需要调用v_target_net.set_weights(v_main_net.get_weights())即可,非常直观。
内容的提问来源于stack exchange,提问作者yjc
相关产品推荐
相关产品推荐

