TensorFlow中如何共享LSTMCell、Dense等类对象的权重参数?
如何在TensorFlow中复用LSTMCell/Dense等层的权重
当然可以实现!其实tf.nn.rnn_cell.LSTMCell、tf.layers.Dense这类层的核心权重本质上还是tf.Variable,所以我们可以借助变量作用域(variable_scope)的复用机制,或者直接复用已创建的层对象,来实现训练和预测阶段的权重共享。下面分两种常用方法详细说明:
方法1:利用variable_scope的复用模式
这类层内部都是通过tf.get_variable()来创建权重变量的,只要我们在创建层时指定统一的变量作用域,然后在预测阶段开启复用模式,新创建的层就会自动复用之前的权重。
示例:LSTMCell的复用
# 训练阶段:创建带指定作用域的LSTMCell with tf.variable_scope("my_lstm"): lstm_cell_1 = tf.nn.rnn_cell.LSTMCell(num_units=256) # 执行训练前向传播 train_outputs, _ = tf.nn.dynamic_rnn(lstm_cell_1, train_inputs, dtype=tf.float32) # 预测阶段:复用同一个作用域的变量 with tf.variable_scope("my_lstm", reuse=True): lstm_cell_2 = tf.nn.rnn_cell.LSTMCell(num_units=256) # 这里lstm_cell_2会完全复用lstm_cell_1的权重和偏置 pred_outputs, _ = tf.nn.dynamic_rnn(lstm_cell_2, pred_inputs, dtype=tf.float32)
示例:Dense层的复用
# 训练阶段 with tf.variable_scope("my_dense"): dense_1 = tf.layers.Dense(units=128, activation=tf.nn.relu) train_logits = dense_1(train_outputs[:, -1, :]) # 预测阶段 with tf.variable_scope("my_dense", reuse=True): dense_2 = tf.layers.Dense(units=128, activation=tf.nn.relu) pred_logits = dense_2(pred_outputs[:, -1, :])
方法2:直接复用已创建的层对象
更简单的方式是,训练阶段创建的层对象本身就可以直接在预测阶段使用——因为层对象会自己管理内部的变量,只要在同一个计算图中,直接调用它的前向传播方法即可,完全不需要重新创建新的层对象。
# 一次性创建训练和预测都要用的层 lstm_cell = tf.nn.rnn_cell.LSTMCell(num_units=256) dense_layer = tf.layers.Dense(units=vocab_size) # 语言建模的输出层,对应词汇表大小 # 训练阶段前向传播 train_outputs, _ = tf.nn.dynamic_rnn(lstm_cell, train_inputs, dtype=tf.float32) train_logits = dense_layer(train_outputs[:, -1, :]) # 预测阶段直接用同一个层对象 pred_outputs, _ = tf.nn.dynamic_rnn(lstm_cell, pred_inputs, dtype=tf.float32) pred_logits = dense_layer(pred_outputs[:, -1, :])
额外注意事项
- 如果训练和预测是在不同的会话或者计算图中,你需要先通过
tf.train.Saver()保存训练好的变量,然后在预测时加载这些变量。只要预测阶段创建的层和训练阶段的变量名一致(比如作用域相同、层参数一致),加载后就会自动复用权重。 - 确保训练和预测时的层参数完全一致(比如LSTM的
num_units、Dense的units),否则会因为变量形状不匹配导致错误。
内容的提问来源于stack exchange,提问作者user3839908
相关产品推荐
相关产品推荐

