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

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_encoder is tied via tf.transpose(W_encoder) for the decoder. Here, W_xh_enc (encoder input→hidden) is tied to W_hy_dec (decoder hidden→output) via transpose—this directly mirrors the fully-connected weight binding.
  • The recurrent weight W_hh_enc is also tied via transpose to W_hh_dec to 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 RNNCell so you can see exactly where the tied weights are used—no black boxes!

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.21 03:49:23