You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

向LSTMCell输入初始状态:GRU转LSTM及RNN Dropout应用疑问

Hey there! Let's tackle your two questions about switching from GRU to LSTM cells and properly applying dropout in your TensorFlow RNN.

1. Feeding Initial State to LSTMCell

Unlike GRU cells (which only have a single hidden state), LSTM cells maintain a tuple of two states: the cell state (c, which holds long-term memory) and the hidden state (h, which is the output state). This means you need to adjust your initial state handling compared to the GRU setup.

Here's how to adapt your code:

Step 1: Define the right initial state placeholders

You have two common approaches here:

  • Option 1: Split into separate placeholders for cell and hidden states

    # For a multi-layer LSTM (NLAYERS=3), each layer has its own c and h
    # Total size per state type: INTERNALSIZE * NLAYERS
    Hin_c = tf.placeholder(tf.float32, [None, INTERNALSIZE * NLAYERS], name='Hin_c')
    Hin_h = tf.placeholder(tf.float32, [None, INTERNALSIZE * NLAYERS], name='Hin_h')
    
    # Wrap into LSTMStateTuple for each layer
    initial_state = tuple(
        tf.nn.rnn_cell.LSTMStateTuple(
            Hin_c[:, i*INTERNALSIZE:(i+1)*INTERNALSIZE],
            Hin_h[:, i*INTERNALSIZE:(i+1)*INTERNALSIZE]
        ) for i in range(NLAYERS)
    )
    
  • Option 2: Use a single placeholder and split it
    If you prefer keeping a single placeholder like your original GRU code, you can combine both states into one tensor and split it:

    # Combined size: 2 * INTERNALSIZE * NLAYERS (since each layer has 2 states)
    Hin = tf.placeholder(tf.float32, [None, 2 * INTERNALSIZE * NLAYERS], name='Hin')
    
    # Split into cell and hidden states for each layer
    initial_state = tuple(
        tf.nn.rnn_cell.LSTMStateTuple(
            Hin[:, i*2*INTERNALSIZE : (i*2+1)*INTERNALSIZE],
            Hin[:, (i*2+1)*INTERNALSIZE : (i+1)*2*INTERNALSIZE]
        ) for i in range(NLAYERS)
    )
    

Step 2: Pass initial state to dynamic_rnn

Just like with GRU, you pass the initial_state parameter to tf.nn.dynamic_rnn:

# Define your multi-layer LSTM cell first
lstm_cell = tf.nn.rnn_cell.LSTMCell(INTERNALSIZE)
multi_lstm_cell = tf.nn.rnn_cell.MultiRNNCell([lstm_cell]*NLAYERS)

# Run dynamic_rnn with the initial state
outputs, final_state = tf.nn.dynamic_rnn(
    multi_lstm_cell,
    Xo,
    initial_state=initial_state,
    dtype=tf.float32
)
2. Correctly Applying Dropout in RNNs

Regular dropout (like tf.nn.dropout) isn't ideal for RNNs because it would apply random masks independently at each time step, which breaks the temporal consistency. Instead, use TensorFlow's DropoutWrapper to apply dropout consistently across time steps.

How to implement it:

  1. Wrap each LSTM cell with DropoutWrapper
    Create a helper function to build dropout-wrapped LSTM cells, then stack them into a multi-layer cell:

    # Use a placeholder to toggle dropout between training/testing
    keep_prob = tf.placeholder(tf.float32, name='keep_prob')
    
    def build_dropout_lstm_cell():
        base_cell = tf.nn.rnn_cell.LSTMCell(INTERNALSIZE)
        # Apply dropout to input and output connections (not state connections)
        dropout_cell = tf.nn.rnn_cell.DropoutWrapper(
            base_cell,
            input_keep_prob=keep_prob,  # Dropout on input to the cell
            output_keep_prob=keep_prob, # Dropout on cell output
            state_keep_prob=1.0         # Keep state connections intact (critical for RNN memory)
        )
        return dropout_cell
    
    # Build multi-layer cell with dropout
    multi_lstm_cell = tf.nn.rnn_cell.MultiRNNCell([build_dropout_lstm_cell() for _ in range(NLAYERS)])
    
  2. Toggle dropout during training vs inference

    • During training: Feed keep_prob=0.8 (or your preferred dropout rate, usually 0.5-0.8)
    • During inference: Feed keep_prob=1.0 to disable dropout completely

Key notes:

  • Never apply dropout to the state connections (state_keep_prob=1.0), as this would disrupt the RNN's ability to retain long-term information.
  • You can adjust input_keep_prob and output_keep_prob independently if needed (e.g., higher dropout on inputs than outputs).

内容的提问来源于stack exchange,提问作者lppier

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.21 08:08:36