TensorFlow循环单元中如何用绑定权重构建递归自编码器?
Tied Weights in Recurrent Autoencoders (Basic RNN Implementation)
Got it, let's walk through how to adapt tied weights from fully-connected autoencoders to a recurrent autoencoder using the most basic RNN cells—no fancy pre-built layers, so you can see exactly how the weight sharing works under the hood.
Core Idea Recap
In your fully-connected example, the decoder uses the transpose of the encoder's weight matrix instead of learning a separate set. For a recurrent autoencoder, we extend this logic to the two key weight matrices in a basic RNN:
- The input-to-hidden weight matrix (
W_xh) of the encoder → its transpose becomes the hidden-to-output weight matrix (W_hy) of the decoder - The hidden-to-hidden weight matrix (
W_hh) of the encoder → its transpose becomes the hidden-to-hidden weight matrix of the decoder (since the decoder is effectively "reversing" the encoder's recurrent step)
We'll manually implement the RNN loops to make the weight binding explicit.
TensorFlow Code Example (Basic RNN)
import tensorflow as tf # -------------------------- # Step 1: Define Shared/Tied Weights # -------------------------- input_dim = 10 # Dimension of your input sequence elements hidden_dim = 20 # Dimension of the RNN hidden state # Encoder weights (we'll tie decoder weights to these) W_xh_enc = tf.Variable(tf.random.normal([input_dim, hidden_dim])) # Input to hidden W_hh_enc = tf.Variable(tf.random.normal([hidden_dim, hidden_dim])) # Hidden to hidden b_h_enc = tf.Variable(tf.zeros([hidden_dim])) # Encoder hidden bias # Decoder biases (separate from encoder, like your fully-connected b_decoder) b_h_dec = tf.Variable(tf.zeros([hidden_dim])) # Decoder hidden bias b_y_dec = tf.Variable(tf.zeros([input_dim])) # Decoder output bias # Tied weights: decoder uses transposes of encoder weights W_hy_dec = tf.transpose(W_xh_enc) # Hidden to output (tied to encoder input->hidden) W_hh_dec = tf.transpose(W_hh_enc) # Hidden to hidden (tied to encoder hidden->hidden) # -------------------------- # Step 2: Encoder Forward Pass # -------------------------- def encoder_rnn(input_sequence): # input_sequence shape: [batch_size, sequence_length, input_dim] batch_size = tf.shape(input_sequence)[0] sequence_length = tf.shape(input_sequence)[1] # Initialize hidden state to zeros h_prev = tf.zeros([batch_size, hidden_dim]) for t in range(sequence_length): x_t = input_sequence[:, t, :] # Get t-th time step input # Basic RNN calculation: h_t = tanh(W_xh*x_t + W_hh*h_prev + b_h) h_curr = tf.tanh(tf.matmul(x_t, W_xh_enc) + tf.matmul(h_prev, W_hh_enc) + b_h_enc) h_prev = h_curr # Return final hidden state as the "latent code" return h_prev # -------------------------- # Step 3: Decoder Forward Pass (Generates Reconstructed Sequence) # -------------------------- def decoder_rnn(latent_code, sequence_length): # latent_code shape: [batch_size, hidden_dim] batch_size = tf.shape(latent_code)[0] # Initialize decoder hidden state to the encoder's final latent code h_prev = latent_code reconstructed_sequence = [] for t in range(sequence_length): # Decoder RNN step: reverse of encoder (uses tied weights) h_curr = tf.tanh(tf.matmul(h_prev, W_hh_dec) + b_h_dec) # Generate output at this time step (uses tied W_hy_dec = W_xh_enc^T) y_t = tf.tanh(tf.matmul(h_curr, W_hy_dec) + b_y_dec) # Append to reconstructed sequence reconstructed_sequence.append(y_t) # Update hidden state for next step h_prev = h_curr # Stack sequence into tensor: [batch_size, sequence_length, input_dim] return tf.stack(reconstructed_sequence, axis=1) # -------------------------- # Step 4: Full Autoencoder Pipeline # -------------------------- # Example input: batch of 32 sequences, each 15 time steps long, input_dim=10 input_data = tf.random.normal([32, 15, input_dim]) # Encode latent = encoder_rnn(input_data) # Decode (reconstruct input sequence) reconstructed = decoder_rnn(latent, sequence_length=15) # Loss function (e.g., MSE between input and reconstruction) loss = tf.reduce_mean(tf.square(input_data - reconstructed))
Key Notes to Connect to Your Fully-Connected Example
- In your original code,
W_encoderis tied viatf.transpose(W_encoder)for the decoder. Here,W_xh_enc(encoder input→hidden) is tied toW_hy_dec(decoder hidden→output) via transpose—this directly mirrors the fully-connected weight binding. - The recurrent weight
W_hh_encis also tied via transpose toW_hh_decto maintain the recurrent structure's weight sharing, which is the extra layer of complexity compared to the fully-connected case. - We manually loop through time steps instead of using TensorFlow's pre-built
RNNCellso you can see exactly where the tied weights are used—no black boxes!
内容的提问来源于stack exchange,提问作者Ray
相关产品推荐
相关产品推荐

